Download registration.py from paradigma-inc/limite-1b-violetto: direct link, hf CLI and curl.
- Browser
- Download file 2.61 kB
-
https://huggingface.co/paradigma-inc/limite-1b-violetto/resolve/main/registration.py
- Command line
-
hf download hf://paradigma-inc/limite-1b-violetto/registration.py
-
curl -L -o registration.py https://huggingface.co/paradigma-inc/limite-1b-violetto/resolve/main/registration.py
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.""" | |
| 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)) | |
| def reverse_op(self) -> "_ConcatenateGQAQKV": | |
| return _ConcatenateGQAQKV(self.dim) | |
| class _ConcatenateGQAQKV(Concatenate): | |
| """Fuse Q/K/V while retaining a GQA-aware reverse save transform.""" | |
| 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"] | |