limite-1b-violetto / registration.py
MisterOss's picture
Add Hugging Face Transformers inference support
b60c6b3 verified
Raw History Blame Contribute Delete
2.61 kB
"""Checkpoint-layout registration for the Hub-hosted Limite implementation."""
import torch
from transformers import Chunk, Concatenate, WeightConverter
from transformers.conversion_mapping import register_checkpoint_conversion_mapping
from .configuration_limite import LimiteConfig
class _SplitGQAQKV(Chunk):
"""Reverse a fused GQA projection without assuming equal Q/K/V sizes."""
@torch.no_grad()
def convert(
self,
input_dict: dict[str, torch.Tensor | list[torch.Tensor]],
source_patterns: list[str],
target_patterns: list[str],
*,
config: LimiteConfig,
**kwargs: object,
) -> dict[str, torch.Tensor]:
del source_patterns, kwargs
value = next(iter(input_dict.values()))
tensor = value[0] if isinstance(value, list) else value
q_size = int(config.num_attention_heads) * int(config.head_dim)
kv_size = int(config.num_key_value_heads) * int(config.head_dim)
expected = q_size + 2 * kv_size
if tensor.shape[self.dim] != expected:
raise ValueError(
"Invalid fused QKV size: "
f"expected {expected}, found {tensor.shape[self.dim]}."
)
chunks = tensor.split((q_size, kv_size, kv_size), dim=self.dim)
return dict(zip(target_patterns, chunks, strict=True))
@property
def reverse_op(self) -> "_ConcatenateGQAQKV":
return _ConcatenateGQAQKV(self.dim)
class _ConcatenateGQAQKV(Concatenate):
"""Fuse Q/K/V while retaining a GQA-aware reverse save transform."""
@property
def reverse_op(self) -> _SplitGQAQKV:
return _SplitGQAQKV(self.dim)
def register_weight_converters() -> None:
"""Register split-checkpoint to fused-runtime transformations."""
register_checkpoint_conversion_mapping(
LimiteConfig.model_type,
[
WeightConverter(
source_patterns=[
"self_attn.q_proj.weight",
"self_attn.k_proj.weight",
"self_attn.v_proj.weight",
],
target_patterns="self_attn.qkv_proj.weight",
operations=[_ConcatenateGQAQKV(dim=0)],
),
WeightConverter(
source_patterns=[
"mlp.gate_proj.weight",
"mlp.up_proj.weight",
],
target_patterns="mlp.gate_up_proj.weight",
operations=[Concatenate(dim=0)],
),
],
overwrite=True,
)
__all__ = ["register_weight_converters"]