Add Hugging Face Transformers inference support
Browse files- add the remote Transformers configuration and modeling implementation
- register reversible split-to-fused checkpoint conversions
- expose the architecture through AutoConfig and AutoModel auto_map entries
- document validated inference backends, cache support, and limitations
- preserve the existing checkpoint, tokenizer, chat template, and vLLM path
- README.md +54 -1
- config.json +5 -0
- configuration_limite.py +505 -0
- contract.py +16 -0
- modeling_limite.py +1042 -0
- registration.py +77 -0
- tokenizer_config.json +1 -0
README.md
CHANGED
|
@@ -59,7 +59,60 @@ VLLM_PLUGINS=limite uv run --locked vllm serve paradigma-inc/limite-1b-violetto
|
|
| 59 |
|
| 60 |
The checkpoint is downloaded from Hugging Face on first use. Keep tensor and pipeline parallel sizes at 1. For full setup instructions and the option to install only the plugin into an existing compatible environment, see the [GitHub README](https://github.com/paradigma-inc/limite-violetto#quickstart).
|
| 61 |
|
| 62 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 63 |
|
| 64 |
## Prompting and intended use
|
| 65 |
|
|
|
|
| 59 |
|
| 60 |
The checkpoint is downloaded from Hugging Face on first use. Keep tensor and pipeline parallel sizes at 1. For full setup instructions and the option to install only the plugin into an existing compatible environment, see the [GitHub README](https://github.com/paradigma-inc/limite-violetto#quickstart).
|
| 61 |
|
| 62 |
+
## Run with Transformers
|
| 63 |
+
|
| 64 |
+
Violetto officially supports inference through the standard Hugging Face Transformers APIs. The custom architecture code is downloaded from this repository, so loading the model requires `trust_remote_code=True`.
|
| 65 |
+
|
| 66 |
+
Validated on an NVIDIA H100 with **Python 3.12 · PyTorch 2.11.0 (CUDA 13.0) · Transformers 5.6.2**. Other version and hardware combinations have not yet been formally qualified. SDPA is the portable default.
|
| 67 |
+
|
| 68 |
+
```python
|
| 69 |
+
import torch
|
| 70 |
+
from transformers import AutoModelForCausalLM, AutoTokenizer
|
| 71 |
+
|
| 72 |
+
model_id = "paradigma-inc/limite-1b-violetto"
|
| 73 |
+
|
| 74 |
+
tokenizer = AutoTokenizer.from_pretrained(model_id)
|
| 75 |
+
model = AutoModelForCausalLM.from_pretrained(
|
| 76 |
+
model_id,
|
| 77 |
+
trust_remote_code=True,
|
| 78 |
+
torch_dtype="auto",
|
| 79 |
+
device_map="auto",
|
| 80 |
+
attn_implementation="sdpa",
|
| 81 |
+
)
|
| 82 |
+
|
| 83 |
+
messages = [{"role": "user", "content": "Solve: If x + 3 = 8, what is x?"}]
|
| 84 |
+
text = tokenizer.apply_chat_template(
|
| 85 |
+
messages,
|
| 86 |
+
tokenize=False,
|
| 87 |
+
add_generation_prompt=True,
|
| 88 |
+
)
|
| 89 |
+
inputs = tokenizer([text], return_tensors="pt").to(model.device)
|
| 90 |
+
|
| 91 |
+
with torch.inference_mode():
|
| 92 |
+
output_ids = model.generate(**inputs, max_new_tokens=256)
|
| 93 |
+
|
| 94 |
+
answer = tokenizer.decode(
|
| 95 |
+
output_ids[0, inputs.input_ids.shape[1]:],
|
| 96 |
+
skip_special_tokens=True,
|
| 97 |
+
)
|
| 98 |
+
print(answer)
|
| 99 |
+
```
|
| 100 |
+
|
| 101 |
+
### Supported Transformers inference paths
|
| 102 |
+
|
| 103 |
+
| Attention backend | Cache | Status |
|
| 104 |
+
| --- | --- | --- |
|
| 105 |
+
| SDPA | `DynamicCache` | Supported |
|
| 106 |
+
| SDPA | mixed global/sliding-window `StaticCache` | Supported |
|
| 107 |
+
| SDPA | `StaticCache` + `torch.compile` | Supported |
|
| 108 |
+
| FlashAttention 2 | `DynamicCache` | Supported |
|
| 109 |
+
| FlashAttention 2 | `StaticCache` | Unsupported; rejected with an explicit error |
|
| 110 |
+
|
| 111 |
+
FlashAttention 2 with `StaticCache` is intentionally rejected because that combination does not produce numerically correct logits for Violetto's hybrid local/global attention layout. Use SDPA with `StaticCache`, or FlashAttention 2 with `DynamicCache`.
|
| 112 |
+
|
| 113 |
+
FlashAttention 2 was validated through Transformers' `kernels-community/flash-attn2` integration with `kernels==0.12.3`. The separately installed native `flash_attn` package has not been independently qualified.
|
| 114 |
+
|
| 115 |
+
Official support currently covers inference. The Transformers implementation is differentiable and exposes the standard causal-language-model loss, but full training, gradient checkpointing, PEFT/LoRA, and distributed-training workflows have not yet been formally validated and are not part of the supported interface.
|
| 116 |
|
| 117 |
## Prompting and intended use
|
| 118 |
|
config.json
CHANGED
|
@@ -2,6 +2,11 @@
|
|
| 2 |
"architectures": [
|
| 3 |
"LimiteForCausalLM"
|
| 4 |
],
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 5 |
"attention_softmax_scale": 0.1,
|
| 6 |
"attn_gate_applied": "per_head_before_o_proj",
|
| 7 |
"attn_gate_channels": 128,
|
|
|
|
| 2 |
"architectures": [
|
| 3 |
"LimiteForCausalLM"
|
| 4 |
],
|
| 5 |
+
"auto_map": {
|
| 6 |
+
"AutoConfig": "configuration_limite.LimiteConfig",
|
| 7 |
+
"AutoModel": "modeling_limite.LimiteModel",
|
| 8 |
+
"AutoModelForCausalLM": "modeling_limite.LimiteForCausalLM"
|
| 9 |
+
},
|
| 10 |
"attention_softmax_scale": 0.1,
|
| 11 |
"attn_gate_applied": "per_head_before_o_proj",
|
| 12 |
"attn_gate_channels": 128,
|
configuration_limite.py
ADDED
|
@@ -0,0 +1,505 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Configuration and validation for Limite checkpoints.
|
| 2 |
+
|
| 3 |
+
Every numerical choice is explicit in ``config.json``. Unsupported values fail
|
| 4 |
+
at load time instead of silently selecting different model behavior.
|
| 5 |
+
|
| 6 |
+
Field semantics are documented in the artifact's own `config_field_notes.json`.
|
| 7 |
+
"""
|
| 8 |
+
|
| 9 |
+
from typing import Any
|
| 10 |
+
|
| 11 |
+
from transformers.configuration_utils import PretrainedConfig
|
| 12 |
+
|
| 13 |
+
from .contract import ARCHITECTURE, BOS_TOKEN_ID, EOS_TOKEN_ID, MODEL_TYPE, PAD_TOKEN_ID
|
| 14 |
+
|
| 15 |
+
REQUIRED_CONFIG_FIELDS: tuple[str, ...] = (
|
| 16 |
+
"architectures",
|
| 17 |
+
"model_type",
|
| 18 |
+
"hidden_size",
|
| 19 |
+
"num_hidden_layers",
|
| 20 |
+
"num_attention_heads",
|
| 21 |
+
"num_key_value_heads",
|
| 22 |
+
"head_dim",
|
| 23 |
+
"intermediate_size",
|
| 24 |
+
"vocab_size",
|
| 25 |
+
"padded_vocab_size",
|
| 26 |
+
"tokenizer_vocab_size",
|
| 27 |
+
"max_position_embeddings",
|
| 28 |
+
"tie_word_embeddings",
|
| 29 |
+
"torch_dtype",
|
| 30 |
+
"rms_norm_has_weight",
|
| 31 |
+
"rms_norm_eps_mode",
|
| 32 |
+
"qk_norm",
|
| 33 |
+
"attention_softmax_scale",
|
| 34 |
+
"sliding_window",
|
| 35 |
+
"sliding_window_convention",
|
| 36 |
+
"global_window",
|
| 37 |
+
"global_layers",
|
| 38 |
+
"global_every",
|
| 39 |
+
"global_nope",
|
| 40 |
+
"attn_gate_channels",
|
| 41 |
+
"attn_gate_scale",
|
| 42 |
+
"attn_gate_applied",
|
| 43 |
+
"pos_mode",
|
| 44 |
+
"rope_frac",
|
| 45 |
+
"rope_base_local",
|
| 46 |
+
"rope_base_global",
|
| 47 |
+
"rope_per_layer",
|
| 48 |
+
"rope_n_pairs",
|
| 49 |
+
"rope_style",
|
| 50 |
+
"rope_cos_sin_dtype",
|
| 51 |
+
"ve_dim",
|
| 52 |
+
"ve_layers",
|
| 53 |
+
"ve_gate_channels",
|
| 54 |
+
"ve_gate_scale",
|
| 55 |
+
"ve_head_slice",
|
| 56 |
+
"ve_stored_heads",
|
| 57 |
+
"ve_applied_before_qk_norm",
|
| 58 |
+
"xsa",
|
| 59 |
+
"xsa_layers",
|
| 60 |
+
"xsa_normalize_eps",
|
| 61 |
+
"mudd",
|
| 62 |
+
"mudd_at",
|
| 63 |
+
"mudd_layers",
|
| 64 |
+
"mudd_taps",
|
| 65 |
+
"mudd_inter",
|
| 66 |
+
"mudd_tap_idx",
|
| 67 |
+
"mudd_mlp",
|
| 68 |
+
"mudd_hist_convention",
|
| 69 |
+
"mudd_accumulation",
|
| 70 |
+
"mlp_type",
|
| 71 |
+
"mlp_formula",
|
| 72 |
+
"mlp_ratio",
|
| 73 |
+
"softcap_logits",
|
| 74 |
+
"final_softcap",
|
| 75 |
+
"lm_head_precision_mode",
|
| 76 |
+
"bos_token_id",
|
| 77 |
+
"eos_token_id",
|
| 78 |
+
"pad_token_id",
|
| 79 |
+
"source_format",
|
| 80 |
+
"checkpoint_step",
|
| 81 |
+
"checkpoint_world_size",
|
| 82 |
+
)
|
| 83 |
+
|
| 84 |
+
HEAD_PRECISION_MODES: tuple[str, ...] = ("oracle_exact", "fp32_accumulate")
|
| 85 |
+
|
| 86 |
+
#: Values the modeling code implements. Anything else must fail loudly: every
|
| 87 |
+
#: entry here is a fork in the numerics, and a silent fallback would produce a
|
| 88 |
+
#: plausible-looking model with wrong numbers.
|
| 89 |
+
SUPPORTED = {
|
| 90 |
+
"mlp_type": {"swiglu"},
|
| 91 |
+
"pos_mode": {"rope"},
|
| 92 |
+
"qk_norm": {"rms_pre_rope"},
|
| 93 |
+
"rms_norm_eps_mode": {"torch_finfo_default"},
|
| 94 |
+
"rope_style": {"interleaved_pairs_odd_lane_sign_flip"},
|
| 95 |
+
"rope_cos_sin_dtype": {"bfloat16"},
|
| 96 |
+
"mudd_accumulation": {"ordered_left_to_right"},
|
| 97 |
+
"ve_head_slice": {"first_num_key_value_heads"},
|
| 98 |
+
"lm_head_precision_mode": set(HEAD_PRECISION_MODES),
|
| 99 |
+
"softcap_kind": {"sigmoid"},
|
| 100 |
+
}
|
| 101 |
+
|
| 102 |
+
MLP_FORMULAS = {
|
| 103 |
+
"swiglu": (
|
| 104 |
+
"(silu(gate_proj(h)) * up_proj(h)) @ down_proj, no clamp on either factor"
|
| 105 |
+
),
|
| 106 |
+
}
|
| 107 |
+
|
| 108 |
+
#: The reference applies `k >= q - sliding_window`, i.e. the window is inclusive
|
| 109 |
+
#: of the query token, so a local layer sees `sliding_window + 1` keys. Matching
|
| 110 |
+
#: the raw number against an exclusive-convention kernel silently drops the
|
| 111 |
+
#: oldest key on every local layer.
|
| 112 |
+
SLIDING_WINDOW_CONVENTION = (
|
| 113 |
+
"k >= q - sliding_window, inclusive of the query token "
|
| 114 |
+
"(span = sliding_window + 1 keys)"
|
| 115 |
+
)
|
| 116 |
+
|
| 117 |
+
|
| 118 |
+
def _derive_layer_types(num_hidden_layers: int, global_layers: list[int]) -> list[str]:
|
| 119 |
+
global_layer_set = set(global_layers)
|
| 120 |
+
return [
|
| 121 |
+
"full_attention" if layer_idx in global_layer_set else "sliding_attention"
|
| 122 |
+
for layer_idx in range(num_hidden_layers)
|
| 123 |
+
]
|
| 124 |
+
|
| 125 |
+
|
| 126 |
+
class LimiteConfig(PretrainedConfig):
|
| 127 |
+
"""Transformers configuration for Limite models."""
|
| 128 |
+
|
| 129 |
+
model_type = MODEL_TYPE
|
| 130 |
+
keys_to_ignore_at_inference = ["past_key_values"]
|
| 131 |
+
|
| 132 |
+
base_model_tp_plan = {
|
| 133 |
+
"layers.*.self_attn.qkv_proj": "colwise",
|
| 134 |
+
"layers.*.self_attn.o_proj": "rowwise",
|
| 135 |
+
"layers.*.mlp.gate_up_proj": "packed_colwise",
|
| 136 |
+
"layers.*.mlp.down_proj": "rowwise",
|
| 137 |
+
}
|
| 138 |
+
base_model_pp_plan = {
|
| 139 |
+
"embed_tokens": (["input_ids"], ["inputs_embeds"]),
|
| 140 |
+
"layers": (["hidden_states"], ["hidden_states"]),
|
| 141 |
+
"norm": (["hidden_states"], ["hidden_states"]),
|
| 142 |
+
}
|
| 143 |
+
|
| 144 |
+
@classmethod
|
| 145 |
+
def from_dict(cls, config_dict: dict[str, Any], **kwargs: Any) -> "LimiteConfig":
|
| 146 |
+
missing = [
|
| 147 |
+
field
|
| 148 |
+
for field in REQUIRED_CONFIG_FIELDS
|
| 149 |
+
if (
|
| 150 |
+
config_dict.get("torch_dtype", config_dict.get("dtype"))
|
| 151 |
+
if field == "torch_dtype"
|
| 152 |
+
else config_dict.get(field)
|
| 153 |
+
)
|
| 154 |
+
is None
|
| 155 |
+
]
|
| 156 |
+
if missing:
|
| 157 |
+
raise ValueError(
|
| 158 |
+
f"Limite config is missing required fields {missing}. "
|
| 159 |
+
"Re-export the checkpoint "
|
| 160 |
+
"instead of guessing numerical choices."
|
| 161 |
+
)
|
| 162 |
+
return super().from_dict(config_dict, **kwargs)
|
| 163 |
+
|
| 164 |
+
def __init__(
|
| 165 |
+
self,
|
| 166 |
+
vocab_size: int = 151680,
|
| 167 |
+
padded_vocab_size: int | None = None,
|
| 168 |
+
tokenizer_vocab_size: int | None = None,
|
| 169 |
+
hidden_size: int = 1280,
|
| 170 |
+
intermediate_size: int = 5120,
|
| 171 |
+
num_hidden_layers: int = 48,
|
| 172 |
+
num_attention_heads: int = 10,
|
| 173 |
+
num_key_value_heads: int = 2,
|
| 174 |
+
head_dim: int = 128,
|
| 175 |
+
mlp_ratio: int = 4,
|
| 176 |
+
mlp_type: str = "swiglu",
|
| 177 |
+
mlp_formula: str | None = None,
|
| 178 |
+
max_position_embeddings: int = 8192,
|
| 179 |
+
tie_word_embeddings: bool = True,
|
| 180 |
+
rms_norm_has_weight: bool = False,
|
| 181 |
+
rms_norm_eps_mode: str = "torch_finfo_default",
|
| 182 |
+
qk_norm: str = "rms_pre_rope",
|
| 183 |
+
attention_softmax_scale: float = 0.1,
|
| 184 |
+
sliding_window: int = 1024,
|
| 185 |
+
sliding_window_convention: str = SLIDING_WINDOW_CONVENTION,
|
| 186 |
+
global_window: int = -1,
|
| 187 |
+
global_layers: list[int] | None = None,
|
| 188 |
+
global_every: int = 4,
|
| 189 |
+
global_nope: bool = True,
|
| 190 |
+
attn_gate_channels: int = 0,
|
| 191 |
+
attn_gate_scale: float = 2.0,
|
| 192 |
+
attn_gate_applied: str = "per_head_before_o_proj",
|
| 193 |
+
attention_dropout: float = 0.0,
|
| 194 |
+
pos_mode: str = "rope",
|
| 195 |
+
rope_frac: float = 0.5,
|
| 196 |
+
rope_base_local: float = 1024.0,
|
| 197 |
+
rope_base_global: float = 1024.0,
|
| 198 |
+
rope_per_layer: bool = False,
|
| 199 |
+
rope_n_pairs: int | None = None,
|
| 200 |
+
rope_style: str = "interleaved_pairs_odd_lane_sign_flip",
|
| 201 |
+
rope_cos_sin_dtype: str = "bfloat16",
|
| 202 |
+
ve_dim: int = 128,
|
| 203 |
+
ve_layers: list[int] | None = None,
|
| 204 |
+
ve_gate_channels: int = 12,
|
| 205 |
+
ve_gate_scale: float = 2.0,
|
| 206 |
+
ve_head_slice: str = "first_num_key_value_heads",
|
| 207 |
+
ve_stored_heads: int | None = None,
|
| 208 |
+
ve_applied_before_qk_norm: bool = True,
|
| 209 |
+
xsa: bool = True,
|
| 210 |
+
xsa_layers: list[int] | None = None,
|
| 211 |
+
xsa_normalize_eps: float = 1e-4,
|
| 212 |
+
mudd: bool = True,
|
| 213 |
+
mudd_at: list[int] | None = None,
|
| 214 |
+
mudd_layers: list[int] | None = None,
|
| 215 |
+
mudd_taps: int = 3,
|
| 216 |
+
mudd_inter: int = 32,
|
| 217 |
+
mudd_tap_idx: dict[str, list[int]] | None = None,
|
| 218 |
+
mudd_hist_convention: str = (
|
| 219 |
+
"hist[0] = rms_norm(embedding); hist[j] = output of layer j-1"
|
| 220 |
+
),
|
| 221 |
+
mudd_accumulation: str = "ordered_left_to_right",
|
| 222 |
+
mudd_mlp: bool = False,
|
| 223 |
+
mudd_r_site: str | None = None,
|
| 224 |
+
softcap_logits: dict[str, Any] | None = None,
|
| 225 |
+
final_softcap: float = 0.0,
|
| 226 |
+
lm_head_precision_mode: str = "oracle_exact",
|
| 227 |
+
lm_head_compute_dtype: str | None = None,
|
| 228 |
+
pad_token_id: int | None = PAD_TOKEN_ID,
|
| 229 |
+
bos_token_id: int | None = BOS_TOKEN_ID,
|
| 230 |
+
eos_token_id: int | list[int] | None = EOS_TOKEN_ID,
|
| 231 |
+
source_format: str | None = None,
|
| 232 |
+
source_metadata_assumptions: list[str] | None = None,
|
| 233 |
+
checkpoint_step: int | None = None,
|
| 234 |
+
checkpoint_world_size: int | None = None,
|
| 235 |
+
**kwargs: Any,
|
| 236 |
+
):
|
| 237 |
+
kwargs.setdefault("architectures", [ARCHITECTURE])
|
| 238 |
+
kwargs.setdefault("attn_implementation", "sdpa")
|
| 239 |
+
self.vocab_size = vocab_size
|
| 240 |
+
self.padded_vocab_size = (
|
| 241 |
+
padded_vocab_size if padded_vocab_size is not None else vocab_size
|
| 242 |
+
)
|
| 243 |
+
self.tokenizer_vocab_size = tokenizer_vocab_size
|
| 244 |
+
self.hidden_size = hidden_size
|
| 245 |
+
self.intermediate_size = intermediate_size
|
| 246 |
+
self.num_hidden_layers = num_hidden_layers
|
| 247 |
+
self.num_attention_heads = num_attention_heads
|
| 248 |
+
self.num_key_value_heads = num_key_value_heads
|
| 249 |
+
self.head_dim = head_dim
|
| 250 |
+
self.mlp_ratio = mlp_ratio
|
| 251 |
+
self.mlp_type = mlp_type
|
| 252 |
+
self.mlp_formula = (
|
| 253 |
+
mlp_formula if mlp_formula is not None else MLP_FORMULAS.get(mlp_type)
|
| 254 |
+
)
|
| 255 |
+
self.max_position_embeddings = max_position_embeddings
|
| 256 |
+
self.rms_norm_has_weight = rms_norm_has_weight
|
| 257 |
+
self.rms_norm_eps_mode = rms_norm_eps_mode
|
| 258 |
+
self.qk_norm = qk_norm
|
| 259 |
+
self.attention_softmax_scale = attention_softmax_scale
|
| 260 |
+
self._serialized_sliding_window = int(sliding_window)
|
| 261 |
+
self.sliding_window = self._serialized_sliding_window + 1
|
| 262 |
+
self.sliding_window_convention = sliding_window_convention
|
| 263 |
+
self.global_window = global_window
|
| 264 |
+
self.global_every = global_every
|
| 265 |
+
self.global_layers = sorted(int(x) for x in (global_layers or []))
|
| 266 |
+
self.global_nope = global_nope
|
| 267 |
+
self.attn_gate_channels = attn_gate_channels
|
| 268 |
+
self.attn_gate_scale = attn_gate_scale
|
| 269 |
+
self.attn_gate_applied = attn_gate_applied
|
| 270 |
+
self.attention_dropout = attention_dropout
|
| 271 |
+
self.pos_mode = pos_mode
|
| 272 |
+
self.rope_frac = rope_frac
|
| 273 |
+
self.rope_base_local = rope_base_local
|
| 274 |
+
self.rope_base_global = rope_base_global
|
| 275 |
+
self.rope_per_layer = rope_per_layer
|
| 276 |
+
self.rope_n_pairs = (
|
| 277 |
+
rope_n_pairs
|
| 278 |
+
if rope_n_pairs is not None
|
| 279 |
+
else max(1, int(head_dim * rope_frac) // 2)
|
| 280 |
+
)
|
| 281 |
+
self.rope_style = rope_style
|
| 282 |
+
self.rope_cos_sin_dtype = rope_cos_sin_dtype
|
| 283 |
+
self.ve_dim = ve_dim
|
| 284 |
+
self.ve_layers = sorted(int(x) for x in (ve_layers or []))
|
| 285 |
+
self.ve_gate_channels = ve_gate_channels
|
| 286 |
+
self.ve_gate_scale = ve_gate_scale
|
| 287 |
+
self.ve_head_slice = ve_head_slice
|
| 288 |
+
self.ve_stored_heads = (
|
| 289 |
+
ve_stored_heads if ve_stored_heads is not None else num_attention_heads
|
| 290 |
+
)
|
| 291 |
+
self.ve_applied_before_qk_norm = ve_applied_before_qk_norm
|
| 292 |
+
self.xsa = xsa
|
| 293 |
+
self.xsa_layers = sorted(int(x) for x in (xsa_layers or []))
|
| 294 |
+
self.xsa_normalize_eps = xsa_normalize_eps
|
| 295 |
+
self.mudd = mudd
|
| 296 |
+
self.mudd_at = sorted(int(x) for x in (mudd_at or []))
|
| 297 |
+
self.mudd_layers = sorted(
|
| 298 |
+
int(x) for x in (mudd_layers if mudd_layers is not None else self.mudd_at)
|
| 299 |
+
)
|
| 300 |
+
self.mudd_taps = mudd_taps
|
| 301 |
+
self.mudd_inter = mudd_inter
|
| 302 |
+
self.mudd_tap_idx = {
|
| 303 |
+
str(k): [int(i) for i in v] for k, v in (mudd_tap_idx or {}).items()
|
| 304 |
+
}
|
| 305 |
+
self.mudd_hist_convention = mudd_hist_convention
|
| 306 |
+
self.mudd_accumulation = mudd_accumulation
|
| 307 |
+
self.mudd_mlp = mudd_mlp
|
| 308 |
+
self.mudd_r_site = (mudd_r_site or "resid") if mudd_mlp else None
|
| 309 |
+
self.softcap_logits = dict(
|
| 310 |
+
softcap_logits or {"kind": "sigmoid", "a": 23.0, "b": 5.0, "c": 7.5}
|
| 311 |
+
)
|
| 312 |
+
self.final_softcap = final_softcap
|
| 313 |
+
if lm_head_compute_dtype is not None:
|
| 314 |
+
raise ValueError(
|
| 315 |
+
"Limite config carries the superseded "
|
| 316 |
+
f"lm_head_compute_dtype={lm_head_compute_dtype!r}. That field named "
|
| 317 |
+
"only one of the head's three roundable stages and is not "
|
| 318 |
+
"reinterpreted; re-export the checkpoint."
|
| 319 |
+
)
|
| 320 |
+
self.lm_head_precision_mode = lm_head_precision_mode
|
| 321 |
+
self.source_format = source_format
|
| 322 |
+
self.source_metadata_assumptions = source_metadata_assumptions
|
| 323 |
+
self.checkpoint_step = checkpoint_step
|
| 324 |
+
self.checkpoint_world_size = checkpoint_world_size
|
| 325 |
+
super().__init__(
|
| 326 |
+
pad_token_id=pad_token_id,
|
| 327 |
+
bos_token_id=bos_token_id,
|
| 328 |
+
eos_token_id=eos_token_id,
|
| 329 |
+
tie_word_embeddings=tie_word_embeddings,
|
| 330 |
+
**kwargs,
|
| 331 |
+
)
|
| 332 |
+
if getattr(self, "layer_types", None) is None:
|
| 333 |
+
self.layer_types = _derive_layer_types(
|
| 334 |
+
self.num_hidden_layers,
|
| 335 |
+
self.global_layers,
|
| 336 |
+
)
|
| 337 |
+
self.validate_architecture()
|
| 338 |
+
|
| 339 |
+
def to_dict(self) -> dict[str, Any]:
|
| 340 |
+
"""Serialize the checkpoint convention, not the native cache span."""
|
| 341 |
+
output = super().to_dict()
|
| 342 |
+
output["sliding_window"] = self._serialized_sliding_window
|
| 343 |
+
output.pop("_serialized_sliding_window", None)
|
| 344 |
+
return output
|
| 345 |
+
|
| 346 |
+
@property
|
| 347 |
+
def ve_layer_to_ordinal(self) -> dict[int, int]:
|
| 348 |
+
return {layer: ordinal for ordinal, layer in enumerate(self.ve_layers)}
|
| 349 |
+
|
| 350 |
+
def tap_indices(self, layer_idx: int) -> list[int] | None:
|
| 351 |
+
return self.mudd_tap_idx.get(str(layer_idx))
|
| 352 |
+
|
| 353 |
+
def is_global_layer(self, layer_idx: int) -> bool:
|
| 354 |
+
return layer_idx in set(self.global_layers)
|
| 355 |
+
|
| 356 |
+
def validate_architecture(self) -> None:
|
| 357 |
+
for field, allowed in (
|
| 358 |
+
("mlp_type", SUPPORTED["mlp_type"]),
|
| 359 |
+
("pos_mode", SUPPORTED["pos_mode"]),
|
| 360 |
+
("qk_norm", SUPPORTED["qk_norm"]),
|
| 361 |
+
("rms_norm_eps_mode", SUPPORTED["rms_norm_eps_mode"]),
|
| 362 |
+
("rope_style", SUPPORTED["rope_style"]),
|
| 363 |
+
("rope_cos_sin_dtype", SUPPORTED["rope_cos_sin_dtype"]),
|
| 364 |
+
("mudd_accumulation", SUPPORTED["mudd_accumulation"]),
|
| 365 |
+
("ve_head_slice", SUPPORTED["ve_head_slice"]),
|
| 366 |
+
("lm_head_precision_mode", SUPPORTED["lm_head_precision_mode"]),
|
| 367 |
+
):
|
| 368 |
+
value = getattr(self, field)
|
| 369 |
+
if value not in allowed:
|
| 370 |
+
raise NotImplementedError(
|
| 371 |
+
f"Limite does not implement {field}={value!r}; "
|
| 372 |
+
f"supported: {sorted(allowed)}"
|
| 373 |
+
)
|
| 374 |
+
if self.rms_norm_has_weight:
|
| 375 |
+
raise NotImplementedError(
|
| 376 |
+
"Limite RMS norm carries no learnable gain "
|
| 377 |
+
"(rms_norm_has_weight must be False)."
|
| 378 |
+
)
|
| 379 |
+
if self.sliding_window_convention != SLIDING_WINDOW_CONVENTION:
|
| 380 |
+
raise NotImplementedError(
|
| 381 |
+
f"Unexpected sliding_window_convention "
|
| 382 |
+
f"{self.sliding_window_convention!r}. The window span is an "
|
| 383 |
+
"off-by-one trap; refusing to guess."
|
| 384 |
+
)
|
| 385 |
+
if self.softcap_logits.get("kind") not in SUPPORTED["softcap_kind"]:
|
| 386 |
+
raise NotImplementedError(
|
| 387 |
+
"Limite implements only sigmoid logit softcapping, got "
|
| 388 |
+
f"{self.softcap_logits!r}"
|
| 389 |
+
)
|
| 390 |
+
if self.final_softcap:
|
| 391 |
+
raise NotImplementedError(
|
| 392 |
+
"Limite does not implement the pre-head tanh cap "
|
| 393 |
+
"(final_softcap must be 0)."
|
| 394 |
+
)
|
| 395 |
+
if self.attn_gate_channels < 0:
|
| 396 |
+
raise ValueError("attn_gate_channels must be non-negative.")
|
| 397 |
+
if self.attn_gate_channels:
|
| 398 |
+
if self.attn_gate_scale != 2.0:
|
| 399 |
+
raise NotImplementedError(
|
| 400 |
+
"Limite implements the attention gate only as "
|
| 401 |
+
f"2 * sigmoid(...), got {self.attn_gate_scale}."
|
| 402 |
+
)
|
| 403 |
+
if self.attn_gate_applied != "per_head_before_o_proj":
|
| 404 |
+
raise NotImplementedError(
|
| 405 |
+
"Limite does not implement "
|
| 406 |
+
f"attn_gate_applied={self.attn_gate_applied!r}."
|
| 407 |
+
)
|
| 408 |
+
if not self.global_nope:
|
| 409 |
+
raise NotImplementedError(
|
| 410 |
+
"Limite implements rotary-free global layers only "
|
| 411 |
+
"(global_nope must be True)."
|
| 412 |
+
)
|
| 413 |
+
if self.rope_per_layer and self.rope_base_local != self.rope_base_global:
|
| 414 |
+
raise NotImplementedError(
|
| 415 |
+
"rope_per_layer with distinct bases is unreachable while global "
|
| 416 |
+
"layers skip rotary entirely."
|
| 417 |
+
)
|
| 418 |
+
if self.num_attention_heads % self.num_key_value_heads != 0:
|
| 419 |
+
raise ValueError(
|
| 420 |
+
f"num_attention_heads ({self.num_attention_heads}) must be "
|
| 421 |
+
f"divisible by num_key_value_heads ({self.num_key_value_heads})."
|
| 422 |
+
)
|
| 423 |
+
if self.num_attention_heads * self.head_dim != self.hidden_size:
|
| 424 |
+
raise ValueError(
|
| 425 |
+
"num_attention_heads * head_dim "
|
| 426 |
+
f"({self.num_attention_heads * self.head_dim}) must equal "
|
| 427 |
+
f"hidden_size ({self.hidden_size})."
|
| 428 |
+
)
|
| 429 |
+
if self.mlp_formula != MLP_FORMULAS[self.mlp_type]:
|
| 430 |
+
raise NotImplementedError(
|
| 431 |
+
f"Limite does not implement mlp_formula={self.mlp_formula!r} "
|
| 432 |
+
f"for {self.mlp_type!r}."
|
| 433 |
+
)
|
| 434 |
+
if self.ve_dim > self.head_dim:
|
| 435 |
+
raise ValueError(
|
| 436 |
+
f"ve_dim ({self.ve_dim}) must not exceed head_dim ({self.head_dim})."
|
| 437 |
+
)
|
| 438 |
+
if self.ve_gate_channels > self.hidden_size:
|
| 439 |
+
raise ValueError(
|
| 440 |
+
f"ve_gate_channels ({self.ve_gate_channels}) must not exceed "
|
| 441 |
+
"hidden_size."
|
| 442 |
+
)
|
| 443 |
+
if self.ve_stored_heads not in {
|
| 444 |
+
self.num_attention_heads,
|
| 445 |
+
self.num_key_value_heads,
|
| 446 |
+
}:
|
| 447 |
+
raise ValueError(
|
| 448 |
+
f"ve_stored_heads ({self.ve_stored_heads}) must be query-head "
|
| 449 |
+
f"width ({self.num_attention_heads}) or key-value-head width "
|
| 450 |
+
f"({self.num_key_value_heads})."
|
| 451 |
+
)
|
| 452 |
+
if self.padded_vocab_size != self.vocab_size:
|
| 453 |
+
raise ValueError(
|
| 454 |
+
f"vocab_size ({self.vocab_size}) is the width of the embedding "
|
| 455 |
+
"and head matrices and must equal padded_vocab_size "
|
| 456 |
+
f"({self.padded_vocab_size})."
|
| 457 |
+
)
|
| 458 |
+
for name, layers in (
|
| 459 |
+
("global_layers", self.global_layers),
|
| 460 |
+
("ve_layers", self.ve_layers),
|
| 461 |
+
("xsa_layers", self.xsa_layers),
|
| 462 |
+
("mudd_layers", self.mudd_layers),
|
| 463 |
+
):
|
| 464 |
+
if layers and (layers[0] < 0 or layers[-1] >= self.num_hidden_layers):
|
| 465 |
+
raise ValueError(
|
| 466 |
+
f"{name}={layers} is out of range for "
|
| 467 |
+
f"num_hidden_layers={self.num_hidden_layers}."
|
| 468 |
+
)
|
| 469 |
+
if self.mudd:
|
| 470 |
+
if sorted(int(k) for k in self.mudd_tap_idx) != list(self.mudd_layers):
|
| 471 |
+
raise ValueError(
|
| 472 |
+
f"mudd_tap_idx keys {sorted(self.mudd_tap_idx)} must match "
|
| 473 |
+
f"mudd_layers {self.mudd_layers}."
|
| 474 |
+
)
|
| 475 |
+
for layer, taps in self.mudd_tap_idx.items():
|
| 476 |
+
if len(taps) != self.mudd_taps:
|
| 477 |
+
raise ValueError(
|
| 478 |
+
f"mudd_tap_idx[{layer}] has {len(taps)} taps, expected "
|
| 479 |
+
f"{self.mudd_taps}."
|
| 480 |
+
)
|
| 481 |
+
if max(taps) > int(layer):
|
| 482 |
+
raise ValueError(
|
| 483 |
+
f"mudd_tap_idx[{layer}]={taps} reads a history entry "
|
| 484 |
+
f"that does not exist yet at layer {layer}."
|
| 485 |
+
)
|
| 486 |
+
elif self.mudd_layers:
|
| 487 |
+
raise ValueError("mudd is disabled but mudd_layers is non-empty.")
|
| 488 |
+
if self.mudd_mlp:
|
| 489 |
+
if not self.mudd:
|
| 490 |
+
raise ValueError("mudd_mlp requires the shared MUDD H-way mixer.")
|
| 491 |
+
if self.mudd_r_site != "resid":
|
| 492 |
+
raise NotImplementedError(
|
| 493 |
+
f"Limite implements MUDD R only at the residual base, got "
|
| 494 |
+
f"{self.mudd_r_site!r}."
|
| 495 |
+
)
|
| 496 |
+
if not self.xsa and self.xsa_layers:
|
| 497 |
+
raise ValueError("xsa is disabled but xsa_layers is non-empty.")
|
| 498 |
+
|
| 499 |
+
|
| 500 |
+
__all__ = [
|
| 501 |
+
"LimiteConfig",
|
| 502 |
+
"MLP_FORMULAS",
|
| 503 |
+
"REQUIRED_CONFIG_FIELDS",
|
| 504 |
+
"SLIDING_WINDOW_CONVENTION",
|
| 505 |
+
]
|
contract.py
ADDED
|
@@ -0,0 +1,16 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Stable names shared by every Limite consumer."""
|
| 2 |
+
|
| 3 |
+
MODEL_TYPE = "limite"
|
| 4 |
+
ARCHITECTURE = "LimiteForCausalLM"
|
| 5 |
+
|
| 6 |
+
BOS_TOKEN_ID = 151643
|
| 7 |
+
EOS_TOKEN_ID = 151645
|
| 8 |
+
PAD_TOKEN_ID = 151643
|
| 9 |
+
|
| 10 |
+
__all__ = [
|
| 11 |
+
"ARCHITECTURE",
|
| 12 |
+
"BOS_TOKEN_ID",
|
| 13 |
+
"EOS_TOKEN_ID",
|
| 14 |
+
"MODEL_TYPE",
|
| 15 |
+
"PAD_TOKEN_ID",
|
| 16 |
+
]
|
modeling_limite.py
ADDED
|
@@ -0,0 +1,1042 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Native Transformers implementation of the Limite causal language model.
|
| 2 |
+
|
| 3 |
+
BF16 SDPA is the portable attention path. The implementation also supports
|
| 4 |
+
Transformers' cache protocol so the same model can be used by ``generate``
|
| 5 |
+
without a separate decoding graph.
|
| 6 |
+
"""
|
| 7 |
+
|
| 8 |
+
from __future__ import annotations
|
| 9 |
+
|
| 10 |
+
from contextlib import nullcontext
|
| 11 |
+
from typing import Any
|
| 12 |
+
|
| 13 |
+
import torch
|
| 14 |
+
import torch.nn.functional as F
|
| 15 |
+
from torch import Tensor, nn
|
| 16 |
+
from torch.nn.attention import SDPBackend, sdpa_kernel
|
| 17 |
+
from transformers.cache_utils import Cache, DynamicCache, StaticCache
|
| 18 |
+
from transformers.generation import GenerationMixin
|
| 19 |
+
from transformers.masking_utils import (
|
| 20 |
+
create_causal_mask,
|
| 21 |
+
create_sliding_window_causal_mask,
|
| 22 |
+
)
|
| 23 |
+
from transformers.modeling_outputs import (
|
| 24 |
+
BaseModelOutputWithPast,
|
| 25 |
+
CausalLMOutputWithPast,
|
| 26 |
+
)
|
| 27 |
+
from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel
|
| 28 |
+
|
| 29 |
+
from .configuration_limite import LimiteConfig
|
| 30 |
+
from .registration import register_weight_converters
|
| 31 |
+
|
| 32 |
+
register_weight_converters()
|
| 33 |
+
|
| 34 |
+
|
| 35 |
+
def _uses_external_flash_attention(attn_implementation: str | None) -> bool:
|
| 36 |
+
"""Identify native and Hub-provided FlashAttention implementations."""
|
| 37 |
+
normalized = str(attn_implementation or "").lower().replace("-", "_")
|
| 38 |
+
return "flash_attention" in normalized or "flash_attn" in normalized
|
| 39 |
+
|
| 40 |
+
|
| 41 |
+
def _reject_external_flash_static_cache(attn_implementation: str | None) -> None:
|
| 42 |
+
if _uses_external_flash_attention(attn_implementation):
|
| 43 |
+
raise ValueError(
|
| 44 |
+
"Limite does not support FlashAttention with StaticCache because "
|
| 45 |
+
"that combination can produce incorrect logits. Use "
|
| 46 |
+
"attn_implementation='sdpa' with StaticCache, or use "
|
| 47 |
+
"DynamicCache with FlashAttention."
|
| 48 |
+
)
|
| 49 |
+
|
| 50 |
+
|
| 51 |
+
def _validate_cache_attention_pair(
|
| 52 |
+
*,
|
| 53 |
+
attn_implementation: str | None,
|
| 54 |
+
past_key_values: Cache | None,
|
| 55 |
+
) -> None:
|
| 56 |
+
if isinstance(past_key_values, StaticCache):
|
| 57 |
+
_reject_external_flash_static_cache(attn_implementation)
|
| 58 |
+
|
| 59 |
+
|
| 60 |
+
def rms_norm(x: Tensor) -> Tensor:
|
| 61 |
+
"""Gain-free RMS norm with PyTorch's dtype-dependent default epsilon."""
|
| 62 |
+
return F.rms_norm(x, (x.size(-1),))
|
| 63 |
+
|
| 64 |
+
|
| 65 |
+
class LimiteRMSNorm(nn.Module):
|
| 66 |
+
def forward(self, hidden_states: Tensor) -> Tensor:
|
| 67 |
+
return rms_norm(hidden_states)
|
| 68 |
+
|
| 69 |
+
|
| 70 |
+
class LimiteRotaryEmbedding(nn.Module):
|
| 71 |
+
"""Checkpoint-exact rotary factors with the static frequencies cached."""
|
| 72 |
+
|
| 73 |
+
def __init__(self, config: LimiteConfig) -> None:
|
| 74 |
+
super().__init__()
|
| 75 |
+
self.rope_base_local = float(config.rope_base_local)
|
| 76 |
+
self.rope_n_pairs = int(config.rope_n_pairs)
|
| 77 |
+
self.head_dim = int(config.head_dim)
|
| 78 |
+
self.register_buffer(
|
| 79 |
+
"frequency",
|
| 80 |
+
self._build_frequency(),
|
| 81 |
+
persistent=False,
|
| 82 |
+
)
|
| 83 |
+
|
| 84 |
+
def _build_frequency(self, device: torch.device | None = None) -> Tensor:
|
| 85 |
+
frequency = (1.0 / self.rope_base_local) ** torch.linspace(
|
| 86 |
+
0,
|
| 87 |
+
1,
|
| 88 |
+
steps=self.rope_n_pairs,
|
| 89 |
+
dtype=torch.float32,
|
| 90 |
+
device="cpu",
|
| 91 |
+
)
|
| 92 |
+
frequency = frequency.repeat_interleave(2)
|
| 93 |
+
frequency = torch.cat(
|
| 94 |
+
[frequency, frequency.new_zeros(self.head_dim - frequency.numel())]
|
| 95 |
+
)
|
| 96 |
+
return frequency if device is None else frequency.to(device=device)
|
| 97 |
+
|
| 98 |
+
def _apply(self, fn: Any, recurse: bool = True) -> "LimiteRotaryEmbedding":
|
| 99 |
+
super()._apply(fn, recurse=recurse)
|
| 100 |
+
# Transformers applies ``dtype=...`` to buffers too. RoPE frequencies
|
| 101 |
+
# are part of Limite's FP32 numerical contract, so reconstruct the
|
| 102 |
+
# derived buffer from the CPU-FP32 formula on the destination device.
|
| 103 |
+
self.frequency = self._build_frequency(device=self.frequency.device)
|
| 104 |
+
return self
|
| 105 |
+
|
| 106 |
+
def forward(self, position_ids: Tensor) -> tuple[Tensor, Tensor]:
|
| 107 |
+
theta = position_ids.to(torch.float32).unsqueeze(-1) * self.frequency
|
| 108 |
+
cosine = theta.cos().to(torch.bfloat16).unsqueeze(-2)
|
| 109 |
+
sine = theta.sin().to(torch.bfloat16)
|
| 110 |
+
sine[..., 1::2] *= -1
|
| 111 |
+
return cosine, sine.unsqueeze(-2)
|
| 112 |
+
|
| 113 |
+
|
| 114 |
+
def apply_rotary(x: Tensor, cosine: Tensor, sine: Tensor) -> Tensor:
|
| 115 |
+
paired = x.view(*x.shape[:-1], x.shape[-1] // 2, 2).flip(-1).view(x.shape)
|
| 116 |
+
return cosine * x + sine * paired
|
| 117 |
+
|
| 118 |
+
|
| 119 |
+
def repeat_kv(hidden_states: Tensor, num_groups: int) -> Tensor:
|
| 120 |
+
"""Expand key/value heads for the eager attention oracle."""
|
| 121 |
+
if num_groups == 1:
|
| 122 |
+
return hidden_states
|
| 123 |
+
batch_size, num_kv_heads, sequence_length, head_dim = hidden_states.shape
|
| 124 |
+
hidden_states = hidden_states[:, :, None, :, :].expand(
|
| 125 |
+
batch_size,
|
| 126 |
+
num_kv_heads,
|
| 127 |
+
num_groups,
|
| 128 |
+
sequence_length,
|
| 129 |
+
head_dim,
|
| 130 |
+
)
|
| 131 |
+
return hidden_states.reshape(
|
| 132 |
+
batch_size,
|
| 133 |
+
num_kv_heads * num_groups,
|
| 134 |
+
sequence_length,
|
| 135 |
+
head_dim,
|
| 136 |
+
)
|
| 137 |
+
|
| 138 |
+
|
| 139 |
+
def eager_attention_forward(
|
| 140 |
+
module: nn.Module,
|
| 141 |
+
query: Tensor,
|
| 142 |
+
key: Tensor,
|
| 143 |
+
value: Tensor,
|
| 144 |
+
attention_mask: Tensor | None,
|
| 145 |
+
scaling: float,
|
| 146 |
+
dropout: float = 0.0,
|
| 147 |
+
**kwargs: Any,
|
| 148 |
+
) -> tuple[Tensor, Tensor]:
|
| 149 |
+
"""Reference attention used for backend parity checks."""
|
| 150 |
+
del kwargs
|
| 151 |
+
key = repeat_kv(key, module.num_key_value_groups)
|
| 152 |
+
value = repeat_kv(value, module.num_key_value_groups)
|
| 153 |
+
weights = torch.matmul(query, key.transpose(2, 3)) * scaling
|
| 154 |
+
if attention_mask is not None:
|
| 155 |
+
weights = weights + attention_mask
|
| 156 |
+
probabilities = F.softmax(weights, dim=-1, dtype=torch.float32).to(query.dtype)
|
| 157 |
+
probabilities = F.dropout(
|
| 158 |
+
probabilities,
|
| 159 |
+
p=dropout,
|
| 160 |
+
training=module.training,
|
| 161 |
+
)
|
| 162 |
+
output = torch.matmul(probabilities, value).transpose(1, 2).contiguous()
|
| 163 |
+
return output, probabilities
|
| 164 |
+
|
| 165 |
+
|
| 166 |
+
def _fp32_parameter(*shape: int, initial: float = 0.0) -> nn.Parameter:
|
| 167 |
+
return nn.Parameter(torch.full(shape, initial, dtype=torch.float32))
|
| 168 |
+
|
| 169 |
+
|
| 170 |
+
class LimiteAttention(nn.Module):
|
| 171 |
+
def __init__(self, config: LimiteConfig, layer_idx: int) -> None:
|
| 172 |
+
super().__init__()
|
| 173 |
+
self.config = config
|
| 174 |
+
self.layer_idx = layer_idx
|
| 175 |
+
self.head_dim = int(config.head_dim)
|
| 176 |
+
self.num_heads = int(config.num_attention_heads)
|
| 177 |
+
self.num_kv_heads = int(config.num_key_value_heads)
|
| 178 |
+
self.num_kv_groups = self.num_heads // self.num_kv_heads
|
| 179 |
+
self.num_key_value_groups = self.num_kv_groups
|
| 180 |
+
self.scaling = float(config.attention_softmax_scale)
|
| 181 |
+
self.attention_dropout = float(config.attention_dropout)
|
| 182 |
+
self.is_causal = True
|
| 183 |
+
self.is_global = config.is_global_layer(layer_idx)
|
| 184 |
+
self.window_span = None if self.is_global else int(config.sliding_window)
|
| 185 |
+
self.applies_rope = not (self.is_global and bool(config.global_nope))
|
| 186 |
+
self.has_ve = layer_idx in set(config.ve_layers)
|
| 187 |
+
self.has_xsa = bool(config.xsa) and layer_idx in set(config.xsa_layers)
|
| 188 |
+
self.attn_gate_channels = int(config.attn_gate_channels)
|
| 189 |
+
|
| 190 |
+
hidden_size = int(config.hidden_size)
|
| 191 |
+
self.q_size = self.num_heads * self.head_dim
|
| 192 |
+
self.kv_size = self.num_kv_heads * self.head_dim
|
| 193 |
+
self.qkv_proj = nn.Linear(
|
| 194 |
+
hidden_size,
|
| 195 |
+
self.q_size + 2 * self.kv_size,
|
| 196 |
+
bias=False,
|
| 197 |
+
)
|
| 198 |
+
self.o_proj = nn.Linear(self.num_heads * self.head_dim, hidden_size, bias=False)
|
| 199 |
+
|
| 200 |
+
self.qkv_scale = _fp32_parameter(initial=1.0)
|
| 201 |
+
self.o_scale = _fp32_parameter(initial=1.0)
|
| 202 |
+
self.register_buffer("_inference_qkv_weight", None, persistent=False)
|
| 203 |
+
self.register_buffer("_inference_o_weight", None, persistent=False)
|
| 204 |
+
self.register_buffer("_inference_xsa_alpha", None, persistent=False)
|
| 205 |
+
self.register_buffer("_inference_ve_gate", None, persistent=False)
|
| 206 |
+
self.register_buffer("_inference_attn_gate", None, persistent=False)
|
| 207 |
+
if self.has_xsa:
|
| 208 |
+
self.xsa_alpha = _fp32_parameter(self.num_heads)
|
| 209 |
+
if self.has_ve:
|
| 210 |
+
self.ve_gate = _fp32_parameter(
|
| 211 |
+
int(config.ve_stored_heads), int(config.ve_gate_channels)
|
| 212 |
+
)
|
| 213 |
+
if self.attn_gate_channels:
|
| 214 |
+
self.attn_gate = _fp32_parameter(self.num_heads, self.attn_gate_channels)
|
| 215 |
+
|
| 216 |
+
@staticmethod
|
| 217 |
+
def _scaled(weight: Tensor, scale: Tensor, dtype: torch.dtype) -> Tensor:
|
| 218 |
+
return (scale.to(torch.float32).view(()) * weight).to(dtype)
|
| 219 |
+
|
| 220 |
+
def train(self, mode: bool = True) -> "LimiteAttention":
|
| 221 |
+
super().train(mode)
|
| 222 |
+
if mode:
|
| 223 |
+
self._inference_qkv_weight = None
|
| 224 |
+
self._inference_o_weight = None
|
| 225 |
+
self._inference_xsa_alpha = None
|
| 226 |
+
self._inference_ve_gate = None
|
| 227 |
+
self._inference_attn_gate = None
|
| 228 |
+
else:
|
| 229 |
+
dtype = self.qkv_proj.weight.dtype
|
| 230 |
+
self._inference_qkv_weight = self._scaled(
|
| 231 |
+
self.qkv_proj.weight,
|
| 232 |
+
self.qkv_scale,
|
| 233 |
+
dtype,
|
| 234 |
+
).detach()
|
| 235 |
+
self._inference_o_weight = self._scaled(
|
| 236 |
+
self.o_proj.weight, self.o_scale, self.o_proj.weight.dtype
|
| 237 |
+
).detach()
|
| 238 |
+
if self.has_xsa:
|
| 239 |
+
self._inference_xsa_alpha = torch.tanh(self.xsa_alpha.float()).detach()
|
| 240 |
+
if self.has_ve:
|
| 241 |
+
self._inference_ve_gate = (
|
| 242 |
+
self.ve_gate[: self.num_kv_heads].to(dtype).detach()
|
| 243 |
+
)
|
| 244 |
+
if self.attn_gate_channels:
|
| 245 |
+
self._inference_attn_gate = self.attn_gate.to(dtype).detach()
|
| 246 |
+
return self
|
| 247 |
+
|
| 248 |
+
def _project_qkv(self, hidden_states: Tensor) -> tuple[Tensor, Tensor, Tensor]:
|
| 249 |
+
sizes = (self.q_size, self.kv_size, self.kv_size)
|
| 250 |
+
if not self.training and self._inference_qkv_weight is not None:
|
| 251 |
+
return F.linear(hidden_states, self._inference_qkv_weight).split(
|
| 252 |
+
sizes, dim=-1
|
| 253 |
+
)
|
| 254 |
+
dtype = hidden_states.dtype
|
| 255 |
+
return F.linear(
|
| 256 |
+
hidden_states,
|
| 257 |
+
self._scaled(self.qkv_proj.weight, self.qkv_scale, dtype),
|
| 258 |
+
).split(
|
| 259 |
+
sizes,
|
| 260 |
+
dim=-1,
|
| 261 |
+
)
|
| 262 |
+
|
| 263 |
+
def _apply_value_embeddings(
|
| 264 |
+
self, hidden_states: Tensor, value_embeds: Tensor, value_states: Tensor
|
| 265 |
+
) -> Tensor:
|
| 266 |
+
gate_weight = (
|
| 267 |
+
self._inference_ve_gate
|
| 268 |
+
if not self.training and self._inference_ve_gate is not None
|
| 269 |
+
else self.ve_gate[: self.num_kv_heads].to(hidden_states.dtype)
|
| 270 |
+
)
|
| 271 |
+
gate = float(self.config.ve_gate_scale) * torch.sigmoid(
|
| 272 |
+
F.linear(hidden_states[..., : gate_weight.size(-1)], gate_weight)
|
| 273 |
+
)
|
| 274 |
+
return value_states + gate.unsqueeze(-1) * value_embeds.to(value_states.dtype)
|
| 275 |
+
|
| 276 |
+
def forward(
|
| 277 |
+
self,
|
| 278 |
+
hidden_states: Tensor,
|
| 279 |
+
value_embeds: Tensor | None,
|
| 280 |
+
cosine: Tensor,
|
| 281 |
+
sine: Tensor,
|
| 282 |
+
attention_mask: Tensor | None,
|
| 283 |
+
past_key_values: Cache | None,
|
| 284 |
+
use_cache: bool,
|
| 285 |
+
output_attentions: bool,
|
| 286 |
+
) -> tuple[Tensor, Tensor | None]:
|
| 287 |
+
batch_size, query_length, _ = hidden_states.shape
|
| 288 |
+
query_states, key_states, value_states = self._project_qkv(hidden_states)
|
| 289 |
+
query_states = query_states.view(
|
| 290 |
+
batch_size, query_length, self.num_heads, self.head_dim
|
| 291 |
+
)
|
| 292 |
+
key_states = key_states.view(
|
| 293 |
+
batch_size, query_length, self.num_kv_heads, self.head_dim
|
| 294 |
+
)
|
| 295 |
+
value_states = value_states.view(
|
| 296 |
+
batch_size, query_length, self.num_kv_heads, self.head_dim
|
| 297 |
+
)
|
| 298 |
+
|
| 299 |
+
if self.has_ve and value_embeds is not None:
|
| 300 |
+
value_states = self._apply_value_embeddings(
|
| 301 |
+
hidden_states, value_embeds, value_states
|
| 302 |
+
)
|
| 303 |
+
current_values = value_states
|
| 304 |
+
query_states, key_states = rms_norm(query_states), rms_norm(key_states)
|
| 305 |
+
if self.applies_rope:
|
| 306 |
+
query_states = apply_rotary(query_states, cosine, sine)
|
| 307 |
+
key_states = apply_rotary(key_states, cosine, sine)
|
| 308 |
+
|
| 309 |
+
if use_cache:
|
| 310 |
+
if past_key_values is None:
|
| 311 |
+
raise ValueError("use_cache=True requires a cache instance")
|
| 312 |
+
cached_keys, cached_values = past_key_values.update(
|
| 313 |
+
key_states.transpose(1, 2),
|
| 314 |
+
value_states.transpose(1, 2),
|
| 315 |
+
self.layer_idx,
|
| 316 |
+
)
|
| 317 |
+
key_states = cached_keys.transpose(1, 2)
|
| 318 |
+
value_states = cached_values.transpose(1, 2)
|
| 319 |
+
if (
|
| 320 |
+
self.config._attn_implementation == "sdpa"
|
| 321 |
+
and attention_mask is None
|
| 322 |
+
and query_length == 1
|
| 323 |
+
):
|
| 324 |
+
flash_decode = (
|
| 325 |
+
query_states.is_cuda
|
| 326 |
+
and query_states.dtype in (torch.float16, torch.bfloat16)
|
| 327 |
+
and torch.cuda.get_device_capability(query_states.device)[0] >= 8
|
| 328 |
+
)
|
| 329 |
+
backend_context = (
|
| 330 |
+
sdpa_kernel(SDPBackend.FLASH_ATTENTION)
|
| 331 |
+
if flash_decode
|
| 332 |
+
else nullcontext()
|
| 333 |
+
)
|
| 334 |
+
with backend_context:
|
| 335 |
+
attention_output = (
|
| 336 |
+
F.scaled_dot_product_attention(
|
| 337 |
+
query_states.transpose(1, 2),
|
| 338 |
+
key_states.transpose(1, 2),
|
| 339 |
+
value_states.transpose(1, 2),
|
| 340 |
+
dropout_p=(
|
| 341 |
+
0.0 if not self.training else self.attention_dropout
|
| 342 |
+
),
|
| 343 |
+
scale=self.scaling,
|
| 344 |
+
is_causal=False,
|
| 345 |
+
enable_gqa=True,
|
| 346 |
+
)
|
| 347 |
+
.transpose(1, 2)
|
| 348 |
+
.contiguous()
|
| 349 |
+
)
|
| 350 |
+
probabilities = None
|
| 351 |
+
else:
|
| 352 |
+
attention_interface = ALL_ATTENTION_FUNCTIONS.get_interface(
|
| 353 |
+
self.config._attn_implementation,
|
| 354 |
+
eager_attention_forward,
|
| 355 |
+
)
|
| 356 |
+
attention_output, probabilities = attention_interface(
|
| 357 |
+
self,
|
| 358 |
+
query_states.transpose(1, 2),
|
| 359 |
+
key_states.transpose(1, 2),
|
| 360 |
+
value_states.transpose(1, 2),
|
| 361 |
+
attention_mask,
|
| 362 |
+
dropout=0.0 if not self.training else self.attention_dropout,
|
| 363 |
+
scaling=self.scaling,
|
| 364 |
+
sliding_window=self.window_span,
|
| 365 |
+
output_attentions=output_attentions,
|
| 366 |
+
)
|
| 367 |
+
if self.has_xsa:
|
| 368 |
+
alpha_values = (
|
| 369 |
+
self._inference_xsa_alpha
|
| 370 |
+
if not self.training and self._inference_xsa_alpha is not None
|
| 371 |
+
else torch.tanh(self.xsa_alpha.float())
|
| 372 |
+
)
|
| 373 |
+
if not self.training:
|
| 374 |
+
grouped_output = attention_output.view(
|
| 375 |
+
batch_size,
|
| 376 |
+
query_length,
|
| 377 |
+
self.num_kv_heads,
|
| 378 |
+
self.num_kv_groups,
|
| 379 |
+
self.head_dim,
|
| 380 |
+
)
|
| 381 |
+
value_direction = F.normalize(
|
| 382 |
+
current_values.float(),
|
| 383 |
+
dim=-1,
|
| 384 |
+
eps=float(self.config.xsa_normalize_eps),
|
| 385 |
+
).unsqueeze(3)
|
| 386 |
+
projection = (grouped_output.float() * value_direction).sum(
|
| 387 |
+
-1, keepdim=True
|
| 388 |
+
)
|
| 389 |
+
alpha = alpha_values.view(
|
| 390 |
+
1, 1, self.num_kv_heads, self.num_kv_groups, 1
|
| 391 |
+
)
|
| 392 |
+
attention_output = (
|
| 393 |
+
grouped_output
|
| 394 |
+
- (alpha * projection * value_direction).to(grouped_output.dtype)
|
| 395 |
+
).reshape(batch_size, query_length, self.num_heads, self.head_dim)
|
| 396 |
+
else:
|
| 397 |
+
value_direction = F.normalize(
|
| 398 |
+
current_values.repeat_interleave(self.num_kv_groups, dim=2).float(),
|
| 399 |
+
dim=-1,
|
| 400 |
+
eps=float(self.config.xsa_normalize_eps),
|
| 401 |
+
)
|
| 402 |
+
projection = (attention_output.float() * value_direction).sum(
|
| 403 |
+
-1, keepdim=True
|
| 404 |
+
)
|
| 405 |
+
alpha = alpha_values.view(1, 1, self.num_heads, 1)
|
| 406 |
+
attention_output = attention_output - (
|
| 407 |
+
alpha * projection * value_direction
|
| 408 |
+
).to(attention_output.dtype)
|
| 409 |
+
if self.attn_gate_channels:
|
| 410 |
+
gate_weight = (
|
| 411 |
+
self._inference_attn_gate
|
| 412 |
+
if not self.training and self._inference_attn_gate is not None
|
| 413 |
+
else self.attn_gate.to(hidden_states.dtype)
|
| 414 |
+
)
|
| 415 |
+
gate = float(self.config.attn_gate_scale) * torch.sigmoid(
|
| 416 |
+
F.linear(
|
| 417 |
+
hidden_states[..., : self.attn_gate_channels],
|
| 418 |
+
gate_weight,
|
| 419 |
+
)
|
| 420 |
+
)
|
| 421 |
+
attention_output = attention_output * gate.to(
|
| 422 |
+
attention_output.dtype
|
| 423 |
+
).unsqueeze(-1)
|
| 424 |
+
|
| 425 |
+
attention_output = attention_output.reshape(batch_size, query_length, -1)
|
| 426 |
+
if not self.training and self._inference_o_weight is not None:
|
| 427 |
+
attention_output = F.linear(attention_output, self._inference_o_weight)
|
| 428 |
+
else:
|
| 429 |
+
attention_output = F.linear(
|
| 430 |
+
attention_output,
|
| 431 |
+
self._scaled(
|
| 432 |
+
self.o_proj.weight,
|
| 433 |
+
self.o_scale,
|
| 434 |
+
attention_output.dtype,
|
| 435 |
+
),
|
| 436 |
+
)
|
| 437 |
+
return attention_output, probabilities if output_attentions else None
|
| 438 |
+
|
| 439 |
+
|
| 440 |
+
class LimiteMLP(nn.Module):
|
| 441 |
+
def __init__(self, config: LimiteConfig) -> None:
|
| 442 |
+
super().__init__()
|
| 443 |
+
hidden_size = int(config.hidden_size)
|
| 444 |
+
intermediate_size = int(config.intermediate_size)
|
| 445 |
+
self.intermediate_size = intermediate_size
|
| 446 |
+
self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False)
|
| 447 |
+
self.gate_up_proj = nn.Linear(
|
| 448 |
+
hidden_size,
|
| 449 |
+
2 * intermediate_size,
|
| 450 |
+
bias=False,
|
| 451 |
+
)
|
| 452 |
+
|
| 453 |
+
def forward(self, hidden_states: Tensor) -> Tensor:
|
| 454 |
+
gate, up = F.linear(
|
| 455 |
+
hidden_states,
|
| 456 |
+
self.gate_up_proj.weight,
|
| 457 |
+
).split(self.intermediate_size, dim=-1)
|
| 458 |
+
activated = F.silu(gate) * up
|
| 459 |
+
return self.down_proj(activated)
|
| 460 |
+
|
| 461 |
+
|
| 462 |
+
class LimiteDecoderLayer(nn.Module):
|
| 463 |
+
def __init__(self, config: LimiteConfig, layer_idx: int) -> None:
|
| 464 |
+
super().__init__()
|
| 465 |
+
self.self_attn = LimiteAttention(config, layer_idx)
|
| 466 |
+
self.mlp = LimiteMLP(config)
|
| 467 |
+
self.resid_lambda_attn = _fp32_parameter(initial=1.0)
|
| 468 |
+
self.post_lambda_attn = _fp32_parameter(initial=1.0)
|
| 469 |
+
self.resid_lambda_mlp = _fp32_parameter(initial=1.0)
|
| 470 |
+
self.post_lambda_mlp = _fp32_parameter(initial=1.0)
|
| 471 |
+
self.register_buffer("_inference_residual_scales", None, persistent=False)
|
| 472 |
+
|
| 473 |
+
def train(self, mode: bool = True) -> "LimiteDecoderLayer":
|
| 474 |
+
super().train(mode)
|
| 475 |
+
if mode:
|
| 476 |
+
self._inference_residual_scales = None
|
| 477 |
+
else:
|
| 478 |
+
dtype = self.self_attn.qkv_proj.weight.dtype
|
| 479 |
+
self._inference_residual_scales = (
|
| 480 |
+
torch.stack(
|
| 481 |
+
(
|
| 482 |
+
self.resid_lambda_attn,
|
| 483 |
+
self.post_lambda_attn,
|
| 484 |
+
self.resid_lambda_mlp,
|
| 485 |
+
self.post_lambda_mlp,
|
| 486 |
+
)
|
| 487 |
+
)
|
| 488 |
+
.to(dtype)
|
| 489 |
+
.detach()
|
| 490 |
+
)
|
| 491 |
+
return self
|
| 492 |
+
|
| 493 |
+
def forward(
|
| 494 |
+
self,
|
| 495 |
+
hidden_states: Tensor,
|
| 496 |
+
attention_input: Tensor,
|
| 497 |
+
value_embeds: Tensor | None,
|
| 498 |
+
cosine: Tensor,
|
| 499 |
+
sine: Tensor,
|
| 500 |
+
attention_mask: Tensor | None,
|
| 501 |
+
past_key_values: Cache | None,
|
| 502 |
+
use_cache: bool,
|
| 503 |
+
output_attentions: bool,
|
| 504 |
+
attention_residual: Tensor | None = None,
|
| 505 |
+
) -> tuple[Tensor, Tensor | None]:
|
| 506 |
+
attention_output, probabilities = self.self_attn(
|
| 507 |
+
attention_input,
|
| 508 |
+
value_embeds,
|
| 509 |
+
cosine,
|
| 510 |
+
sine,
|
| 511 |
+
attention_mask,
|
| 512 |
+
past_key_values,
|
| 513 |
+
use_cache,
|
| 514 |
+
output_attentions,
|
| 515 |
+
)
|
| 516 |
+
residual_base = (
|
| 517 |
+
hidden_states if attention_residual is None else attention_residual
|
| 518 |
+
)
|
| 519 |
+
use_constants = (
|
| 520 |
+
not self.training and self._inference_residual_scales is not None
|
| 521 |
+
)
|
| 522 |
+
if use_constants:
|
| 523 |
+
residual_scales = self._inference_residual_scales
|
| 524 |
+
mixed = (
|
| 525 |
+
residual_scales[0] * residual_base
|
| 526 |
+
+ residual_scales[1] * attention_output
|
| 527 |
+
)
|
| 528 |
+
else:
|
| 529 |
+
mixed = (
|
| 530 |
+
self.resid_lambda_attn.to(hidden_states.dtype) * residual_base
|
| 531 |
+
+ self.post_lambda_attn.to(attention_output.dtype) * attention_output
|
| 532 |
+
)
|
| 533 |
+
mlp_output = self.mlp(rms_norm(mixed))
|
| 534 |
+
if use_constants:
|
| 535 |
+
output = residual_scales[2] * mixed + residual_scales[3] * mlp_output
|
| 536 |
+
else:
|
| 537 |
+
output = (
|
| 538 |
+
self.resid_lambda_mlp.to(mixed.dtype) * mixed
|
| 539 |
+
+ self.post_lambda_mlp.to(mlp_output.dtype) * mlp_output
|
| 540 |
+
)
|
| 541 |
+
return output, probabilities
|
| 542 |
+
|
| 543 |
+
|
| 544 |
+
class LimiteMudd(nn.Module):
|
| 545 |
+
def __init__(self, config: LimiteConfig) -> None:
|
| 546 |
+
super().__init__()
|
| 547 |
+
self.dense1 = _fp32_parameter(int(config.mudd_inter), int(config.hidden_size))
|
| 548 |
+
self.dense2 = _fp32_parameter(
|
| 549 |
+
int(config.num_hidden_layers),
|
| 550 |
+
int(config.mudd_taps),
|
| 551 |
+
int(config.mudd_inter),
|
| 552 |
+
)
|
| 553 |
+
self.bias = _fp32_parameter(
|
| 554 |
+
int(config.num_hidden_layers), int(config.mudd_taps)
|
| 555 |
+
)
|
| 556 |
+
self.register_buffer("_inference_dense1", None, persistent=False)
|
| 557 |
+
self.register_buffer("_inference_dense2", None, persistent=False)
|
| 558 |
+
self.register_buffer("_inference_bias", None, persistent=False)
|
| 559 |
+
self.register_buffer("_inference_dense2_mlp", None, persistent=False)
|
| 560 |
+
self.register_buffer("_inference_bias_mlp", None, persistent=False)
|
| 561 |
+
self.uses_r_way = bool(config.mudd_mlp)
|
| 562 |
+
if self.uses_r_way:
|
| 563 |
+
self.dense2_mlp = _fp32_parameter(
|
| 564 |
+
int(config.num_hidden_layers),
|
| 565 |
+
int(config.mudd_taps),
|
| 566 |
+
int(config.mudd_inter),
|
| 567 |
+
)
|
| 568 |
+
self.bias_mlp = _fp32_parameter(
|
| 569 |
+
int(config.num_hidden_layers), int(config.mudd_taps)
|
| 570 |
+
)
|
| 571 |
+
|
| 572 |
+
def train(self, mode: bool = True) -> "LimiteMudd":
|
| 573 |
+
super().train(mode)
|
| 574 |
+
if mode:
|
| 575 |
+
self._inference_dense1 = None
|
| 576 |
+
self._inference_dense2 = None
|
| 577 |
+
self._inference_bias = None
|
| 578 |
+
self._inference_dense2_mlp = None
|
| 579 |
+
self._inference_bias_mlp = None
|
| 580 |
+
else:
|
| 581 |
+
self._inference_dense1 = self.dense1.to(torch.bfloat16).detach()
|
| 582 |
+
self._inference_dense2 = self.dense2.to(torch.bfloat16).detach()
|
| 583 |
+
self._inference_bias = self.bias.to(torch.bfloat16).detach()
|
| 584 |
+
if self.uses_r_way:
|
| 585 |
+
self._inference_dense2_mlp = self.dense2_mlp.to(torch.bfloat16).detach()
|
| 586 |
+
self._inference_bias_mlp = self.bias_mlp.to(torch.bfloat16).detach()
|
| 587 |
+
return self
|
| 588 |
+
|
| 589 |
+
def _inner(self, current: Tensor) -> Tensor:
|
| 590 |
+
use_constants = not self.training and self._inference_dense1 is not None
|
| 591 |
+
dense1 = (
|
| 592 |
+
self._inference_dense1 if use_constants else self.dense1.to(current.dtype)
|
| 593 |
+
)
|
| 594 |
+
return F.gelu(F.linear(rms_norm(current), dense1))
|
| 595 |
+
|
| 596 |
+
def _mix(
|
| 597 |
+
self,
|
| 598 |
+
taps: list[Tensor],
|
| 599 |
+
inner: Tensor,
|
| 600 |
+
layer_idx: int,
|
| 601 |
+
*,
|
| 602 |
+
r_way: bool,
|
| 603 |
+
) -> Tensor:
|
| 604 |
+
count = len(taps)
|
| 605 |
+
use_constants = not self.training and self._inference_dense1 is not None
|
| 606 |
+
if r_way:
|
| 607 |
+
if not self.uses_r_way:
|
| 608 |
+
raise ValueError("R-way mixing requested without R-way weights")
|
| 609 |
+
if use_constants:
|
| 610 |
+
dense2, bias = (
|
| 611 |
+
self._inference_dense2_mlp,
|
| 612 |
+
self._inference_bias_mlp,
|
| 613 |
+
)
|
| 614 |
+
else:
|
| 615 |
+
dense2, bias = self.dense2_mlp, self.bias_mlp
|
| 616 |
+
elif use_constants:
|
| 617 |
+
dense2, bias = self._inference_dense2, self._inference_bias
|
| 618 |
+
else:
|
| 619 |
+
dense2, bias = self.dense2, self.bias
|
| 620 |
+
weights = torch.einsum(
|
| 621 |
+
"btk,mk->btm",
|
| 622 |
+
inner,
|
| 623 |
+
dense2[layer_idx, :count].to(inner.dtype),
|
| 624 |
+
)
|
| 625 |
+
weights = weights + bias[layer_idx, :count].to(weights.dtype)
|
| 626 |
+
output = weights[..., 0:1].type_as(taps[0]) * taps[0]
|
| 627 |
+
for index in range(1, count):
|
| 628 |
+
output = (
|
| 629 |
+
output
|
| 630 |
+
+ weights[..., index : index + 1].type_as(taps[index]) * taps[index]
|
| 631 |
+
)
|
| 632 |
+
return output
|
| 633 |
+
|
| 634 |
+
def forward(
|
| 635 |
+
self,
|
| 636 |
+
taps: list[Tensor],
|
| 637 |
+
current: Tensor,
|
| 638 |
+
layer_idx: int,
|
| 639 |
+
*,
|
| 640 |
+
r_way: bool = False,
|
| 641 |
+
) -> Tensor:
|
| 642 |
+
return self._mix(
|
| 643 |
+
taps,
|
| 644 |
+
self._inner(current),
|
| 645 |
+
layer_idx,
|
| 646 |
+
r_way=r_way,
|
| 647 |
+
)
|
| 648 |
+
|
| 649 |
+
def forward_pair(
|
| 650 |
+
self,
|
| 651 |
+
taps: list[Tensor],
|
| 652 |
+
current: Tensor,
|
| 653 |
+
layer_idx: int,
|
| 654 |
+
) -> tuple[Tensor, Tensor]:
|
| 655 |
+
if not self.uses_r_way:
|
| 656 |
+
raise ValueError("paired MUDD mixing requires R-way weights")
|
| 657 |
+
inner = self._inner(current)
|
| 658 |
+
return (
|
| 659 |
+
self._mix(taps, inner, layer_idx, r_way=False),
|
| 660 |
+
self._mix(taps, inner, layer_idx, r_way=True),
|
| 661 |
+
)
|
| 662 |
+
|
| 663 |
+
|
| 664 |
+
class LimitePreTrainedModel(PreTrainedModel):
|
| 665 |
+
config_class = LimiteConfig
|
| 666 |
+
base_model_prefix = "model"
|
| 667 |
+
_no_split_modules = ["LimiteDecoderLayer"]
|
| 668 |
+
_skip_keys_device_placement = ["past_key_values"]
|
| 669 |
+
_supports_flash_attn = True
|
| 670 |
+
_supports_sdpa = True
|
| 671 |
+
_supports_attention_backend = True
|
| 672 |
+
_can_compile_fullgraph = True
|
| 673 |
+
|
| 674 |
+
@torch.no_grad()
|
| 675 |
+
def _init_weights(self, module: nn.Module) -> None:
|
| 676 |
+
super()._init_weights(module)
|
| 677 |
+
if isinstance(module, LimiteRotaryEmbedding):
|
| 678 |
+
# Transformers materializes non-persistent buffers with
|
| 679 |
+
# ``empty_like`` during low-memory/device-map loading. Restore this
|
| 680 |
+
# derived FP32 buffer before the loaded model is returned.
|
| 681 |
+
module.frequency.copy_(
|
| 682 |
+
module._build_frequency(device=module.frequency.device)
|
| 683 |
+
)
|
| 684 |
+
|
| 685 |
+
|
| 686 |
+
class LimiteModel(LimitePreTrainedModel):
|
| 687 |
+
def __init__(self, config: LimiteConfig) -> None:
|
| 688 |
+
super().__init__(config)
|
| 689 |
+
self.vocab_size = int(config.vocab_size)
|
| 690 |
+
self.embed_tokens = nn.Embedding(
|
| 691 |
+
int(config.vocab_size), int(config.hidden_size)
|
| 692 |
+
)
|
| 693 |
+
self.value_embeds = nn.Embedding(
|
| 694 |
+
int(config.vocab_size), int(config.ve_stored_heads) * int(config.ve_dim)
|
| 695 |
+
)
|
| 696 |
+
self.layers = nn.ModuleList(
|
| 697 |
+
[
|
| 698 |
+
LimiteDecoderLayer(config, layer_idx)
|
| 699 |
+
for layer_idx in range(int(config.num_hidden_layers))
|
| 700 |
+
]
|
| 701 |
+
)
|
| 702 |
+
self.norm = LimiteRMSNorm()
|
| 703 |
+
self.rotary_emb = LimiteRotaryEmbedding(config)
|
| 704 |
+
self.mudd = LimiteMudd(config) if config.mudd else None
|
| 705 |
+
self.retained_taps = sorted(
|
| 706 |
+
{
|
| 707 |
+
tap
|
| 708 |
+
for layer, taps in config.mudd_tap_idx.items()
|
| 709 |
+
for tap in taps
|
| 710 |
+
if tap != int(layer)
|
| 711 |
+
}
|
| 712 |
+
)
|
| 713 |
+
self.post_init()
|
| 714 |
+
|
| 715 |
+
def get_input_embeddings(self) -> nn.Module:
|
| 716 |
+
return self.embed_tokens
|
| 717 |
+
|
| 718 |
+
def set_input_embeddings(self, value: nn.Module) -> None:
|
| 719 |
+
self.embed_tokens = value
|
| 720 |
+
|
| 721 |
+
def _value_embeddings(self, input_ids: Tensor) -> Tensor | None:
|
| 722 |
+
if not self.config.ve_layers:
|
| 723 |
+
return None
|
| 724 |
+
config = self.config
|
| 725 |
+
embeddings = self.value_embeds(input_ids).view(
|
| 726 |
+
*input_ids.shape, int(config.ve_stored_heads), int(config.ve_dim)
|
| 727 |
+
)
|
| 728 |
+
if config.ve_dim < config.head_dim:
|
| 729 |
+
embeddings = F.pad(
|
| 730 |
+
embeddings,
|
| 731 |
+
(0, int(config.head_dim) - int(config.ve_dim)),
|
| 732 |
+
)
|
| 733 |
+
return embeddings[..., : int(config.num_key_value_heads), :].contiguous()
|
| 734 |
+
|
| 735 |
+
def forward(
|
| 736 |
+
self,
|
| 737 |
+
input_ids: Tensor | None = None,
|
| 738 |
+
attention_mask: Tensor | None = None,
|
| 739 |
+
position_ids: Tensor | None = None,
|
| 740 |
+
past_key_values: Cache | None = None,
|
| 741 |
+
use_cache: bool | None = None,
|
| 742 |
+
inputs_embeds: Tensor | None = None,
|
| 743 |
+
output_attentions: bool | None = None,
|
| 744 |
+
output_hidden_states: bool | None = None,
|
| 745 |
+
return_dict: bool | None = None,
|
| 746 |
+
**kwargs: Any,
|
| 747 |
+
) -> BaseModelOutputWithPast | tuple[Tensor, ...]:
|
| 748 |
+
del kwargs
|
| 749 |
+
if input_ids is None or inputs_embeds is not None:
|
| 750 |
+
raise ValueError(
|
| 751 |
+
"Limite requires input_ids because value embeddings are a "
|
| 752 |
+
"second token lookup"
|
| 753 |
+
)
|
| 754 |
+
use_cache = self.config.use_cache if use_cache is None else use_cache
|
| 755 |
+
output_attentions = bool(output_attentions)
|
| 756 |
+
output_hidden_states = bool(output_hidden_states)
|
| 757 |
+
return_dict = (
|
| 758 |
+
self.config.use_return_dict if return_dict is None else return_dict
|
| 759 |
+
)
|
| 760 |
+
if output_attentions:
|
| 761 |
+
raise ValueError(
|
| 762 |
+
"Limite does not materialize attention weights; "
|
| 763 |
+
"output_attentions=True is unsupported"
|
| 764 |
+
)
|
| 765 |
+
|
| 766 |
+
if use_cache and past_key_values is None:
|
| 767 |
+
past_key_values = DynamicCache(config=self.config)
|
| 768 |
+
if use_cache:
|
| 769 |
+
_validate_cache_attention_pair(
|
| 770 |
+
attn_implementation=self.config._attn_implementation,
|
| 771 |
+
past_key_values=past_key_values,
|
| 772 |
+
)
|
| 773 |
+
if (
|
| 774 |
+
use_cache
|
| 775 |
+
and isinstance(past_key_values, DynamicCache)
|
| 776 |
+
and not hasattr(past_key_values, "_limite_unpadded")
|
| 777 |
+
):
|
| 778 |
+
if attention_mask is None:
|
| 779 |
+
past_key_values._limite_unpadded = True
|
| 780 |
+
elif isinstance(attention_mask, Tensor):
|
| 781 |
+
past_key_values._limite_unpadded = not bool(
|
| 782 |
+
(attention_mask == 0).any().item()
|
| 783 |
+
)
|
| 784 |
+
else:
|
| 785 |
+
past_key_values._limite_unpadded = False
|
| 786 |
+
if position_ids is None:
|
| 787 |
+
if attention_mask is not None:
|
| 788 |
+
position_ids = attention_mask.long().cumsum(-1) - 1
|
| 789 |
+
position_ids.masked_fill_(attention_mask == 0, 0)
|
| 790 |
+
position_ids = position_ids[:, -input_ids.shape[1] :]
|
| 791 |
+
else:
|
| 792 |
+
past_length = (
|
| 793 |
+
past_key_values.get_seq_length()
|
| 794 |
+
if past_key_values is not None
|
| 795 |
+
else 0
|
| 796 |
+
)
|
| 797 |
+
position_ids = (
|
| 798 |
+
torch.arange(
|
| 799 |
+
past_length,
|
| 800 |
+
past_length + input_ids.shape[1],
|
| 801 |
+
device=input_ids.device,
|
| 802 |
+
)
|
| 803 |
+
.unsqueeze(0)
|
| 804 |
+
.expand(input_ids.shape[0], -1)
|
| 805 |
+
)
|
| 806 |
+
|
| 807 |
+
cosine, sine = self.rotary_emb(position_ids)
|
| 808 |
+
value_embeds = self._value_embeddings(input_ids)
|
| 809 |
+
hidden_states = rms_norm(self.embed_tokens(input_ids))
|
| 810 |
+
maskless_sdpa_decode = (
|
| 811 |
+
self.config._attn_implementation == "sdpa"
|
| 812 |
+
and isinstance(past_key_values, DynamicCache)
|
| 813 |
+
and input_ids.shape[1] == 1
|
| 814 |
+
and bool(getattr(past_key_values, "_limite_unpadded", False))
|
| 815 |
+
)
|
| 816 |
+
if isinstance(attention_mask, dict):
|
| 817 |
+
causal_mask_mapping = attention_mask
|
| 818 |
+
elif maskless_sdpa_decode:
|
| 819 |
+
causal_mask_mapping = {
|
| 820 |
+
"full_attention": None,
|
| 821 |
+
"sliding_attention": None,
|
| 822 |
+
}
|
| 823 |
+
else:
|
| 824 |
+
mask_kwargs = {
|
| 825 |
+
"config": self.config,
|
| 826 |
+
"inputs_embeds": hidden_states,
|
| 827 |
+
"attention_mask": attention_mask,
|
| 828 |
+
"past_key_values": past_key_values,
|
| 829 |
+
"position_ids": position_ids,
|
| 830 |
+
}
|
| 831 |
+
causal_mask_mapping = {
|
| 832 |
+
"full_attention": create_causal_mask(**mask_kwargs),
|
| 833 |
+
"sliding_attention": create_sliding_window_causal_mask(**mask_kwargs),
|
| 834 |
+
}
|
| 835 |
+
history: dict[int, Tensor] = (
|
| 836 |
+
{0: hidden_states} if 0 in self.retained_taps else {}
|
| 837 |
+
)
|
| 838 |
+
all_hidden_states: tuple[Tensor, ...] = ()
|
| 839 |
+
all_attentions: tuple[Tensor, ...] = ()
|
| 840 |
+
|
| 841 |
+
for layer_idx, decoder_layer in enumerate(self.layers):
|
| 842 |
+
if output_hidden_states:
|
| 843 |
+
all_hidden_states += (hidden_states,)
|
| 844 |
+
taps = self.config.tap_indices(layer_idx)
|
| 845 |
+
if taps is not None:
|
| 846 |
+
tap_values = [
|
| 847 |
+
history[tap] if tap != layer_idx else hidden_states for tap in taps
|
| 848 |
+
]
|
| 849 |
+
if not self.training and self.mudd.uses_r_way:
|
| 850 |
+
attention_mix, residual_base = self.mudd.forward_pair(
|
| 851 |
+
tap_values, hidden_states, layer_idx
|
| 852 |
+
)
|
| 853 |
+
attention_input = rms_norm(attention_mix)
|
| 854 |
+
else:
|
| 855 |
+
attention_input = rms_norm(
|
| 856 |
+
self.mudd(tap_values, hidden_states, layer_idx)
|
| 857 |
+
)
|
| 858 |
+
residual_base = (
|
| 859 |
+
self.mudd(
|
| 860 |
+
tap_values,
|
| 861 |
+
hidden_states,
|
| 862 |
+
layer_idx,
|
| 863 |
+
r_way=True,
|
| 864 |
+
)
|
| 865 |
+
if self.mudd.uses_r_way
|
| 866 |
+
else hidden_states
|
| 867 |
+
)
|
| 868 |
+
else:
|
| 869 |
+
attention_input = rms_norm(hidden_states)
|
| 870 |
+
residual_base = hidden_states
|
| 871 |
+
hidden_states, probabilities = decoder_layer(
|
| 872 |
+
hidden_states,
|
| 873 |
+
attention_input,
|
| 874 |
+
value_embeds,
|
| 875 |
+
cosine,
|
| 876 |
+
sine,
|
| 877 |
+
causal_mask_mapping[self.config.layer_types[layer_idx]],
|
| 878 |
+
past_key_values,
|
| 879 |
+
use_cache,
|
| 880 |
+
output_attentions,
|
| 881 |
+
residual_base,
|
| 882 |
+
)
|
| 883 |
+
if output_attentions:
|
| 884 |
+
all_attentions += (probabilities,)
|
| 885 |
+
if layer_idx + 1 in self.retained_taps:
|
| 886 |
+
history[layer_idx + 1] = hidden_states
|
| 887 |
+
|
| 888 |
+
hidden_states = self.norm(hidden_states)
|
| 889 |
+
if output_hidden_states:
|
| 890 |
+
all_hidden_states += (hidden_states,)
|
| 891 |
+
if not return_dict:
|
| 892 |
+
values: tuple[Tensor | Cache | tuple[Tensor, ...], ...] = (hidden_states,)
|
| 893 |
+
if use_cache:
|
| 894 |
+
values += (past_key_values,)
|
| 895 |
+
if output_hidden_states:
|
| 896 |
+
values += (all_hidden_states,)
|
| 897 |
+
if output_attentions:
|
| 898 |
+
values += (all_attentions,)
|
| 899 |
+
return values
|
| 900 |
+
return BaseModelOutputWithPast(
|
| 901 |
+
last_hidden_state=hidden_states,
|
| 902 |
+
past_key_values=past_key_values if use_cache else None,
|
| 903 |
+
hidden_states=all_hidden_states if output_hidden_states else None,
|
| 904 |
+
attentions=all_attentions if output_attentions else None,
|
| 905 |
+
)
|
| 906 |
+
|
| 907 |
+
|
| 908 |
+
class LimiteForCausalLM(LimitePreTrainedModel, GenerationMixin):
|
| 909 |
+
_tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"}
|
| 910 |
+
|
| 911 |
+
def __init__(self, config: LimiteConfig) -> None:
|
| 912 |
+
super().__init__(config)
|
| 913 |
+
self.model = LimiteModel(config)
|
| 914 |
+
self.vocab_size = int(config.vocab_size)
|
| 915 |
+
self.lm_head = nn.Linear(
|
| 916 |
+
int(config.hidden_size), int(config.vocab_size), bias=False
|
| 917 |
+
)
|
| 918 |
+
softcap = dict(config.softcap_logits)
|
| 919 |
+
self.softcap_a = float(softcap["a"])
|
| 920 |
+
self.softcap_b = float(softcap["b"])
|
| 921 |
+
self.softcap_c = float(softcap["c"])
|
| 922 |
+
self.head_precision_mode = str(config.lm_head_precision_mode)
|
| 923 |
+
self.post_init()
|
| 924 |
+
|
| 925 |
+
def get_input_embeddings(self) -> nn.Module:
|
| 926 |
+
return self.model.embed_tokens
|
| 927 |
+
|
| 928 |
+
def set_input_embeddings(self, value: nn.Module) -> None:
|
| 929 |
+
self.model.embed_tokens = value
|
| 930 |
+
|
| 931 |
+
def get_output_embeddings(self) -> nn.Module:
|
| 932 |
+
return self.lm_head
|
| 933 |
+
|
| 934 |
+
def set_output_embeddings(self, value: nn.Module) -> None:
|
| 935 |
+
self.lm_head = value
|
| 936 |
+
|
| 937 |
+
def get_decoder(self) -> LimiteModel:
|
| 938 |
+
return self.model
|
| 939 |
+
|
| 940 |
+
def set_decoder(self, decoder: LimiteModel) -> None:
|
| 941 |
+
self.model = decoder
|
| 942 |
+
|
| 943 |
+
def _prepare_static_cache(
|
| 944 |
+
self,
|
| 945 |
+
cache_implementation: str,
|
| 946 |
+
batch_size: int,
|
| 947 |
+
max_cache_len: int,
|
| 948 |
+
model_kwargs: dict[str, Any],
|
| 949 |
+
) -> Cache:
|
| 950 |
+
# GenerationMixin allocates its persistent static cache before the
|
| 951 |
+
# first model forward. Reject the unsupported backend/cache pair here
|
| 952 |
+
# so users receive the Limite contract error instead of failing inside
|
| 953 |
+
# Transformers' cache preparation. SDPA remains entirely native.
|
| 954 |
+
_reject_external_flash_static_cache(self.config._attn_implementation)
|
| 955 |
+
return super()._prepare_static_cache(
|
| 956 |
+
cache_implementation,
|
| 957 |
+
batch_size,
|
| 958 |
+
max_cache_len,
|
| 959 |
+
model_kwargs,
|
| 960 |
+
)
|
| 961 |
+
|
| 962 |
+
def _softcapped_logits(self, hidden_states: Tensor) -> Tensor:
|
| 963 |
+
if self.head_precision_mode == "oracle_exact":
|
| 964 |
+
logits = F.linear(hidden_states, self.lm_head.weight).float()
|
| 965 |
+
elif hidden_states.is_cuda and hidden_states.dtype in (
|
| 966 |
+
torch.bfloat16,
|
| 967 |
+
torch.float16,
|
| 968 |
+
):
|
| 969 |
+
logits = torch.mm(
|
| 970 |
+
hidden_states.reshape(-1, hidden_states.shape[-1]),
|
| 971 |
+
self.lm_head.weight.t(),
|
| 972 |
+
out_dtype=torch.float32,
|
| 973 |
+
).reshape(*hidden_states.shape[:-1], self.lm_head.weight.shape[0])
|
| 974 |
+
else:
|
| 975 |
+
logits = F.linear(hidden_states.float(), self.lm_head.weight.float())
|
| 976 |
+
return self.softcap_a * torch.sigmoid(
|
| 977 |
+
(logits + self.softcap_b) / self.softcap_c
|
| 978 |
+
)
|
| 979 |
+
|
| 980 |
+
def forward(
|
| 981 |
+
self,
|
| 982 |
+
input_ids: Tensor | None = None,
|
| 983 |
+
attention_mask: Tensor | None = None,
|
| 984 |
+
position_ids: Tensor | None = None,
|
| 985 |
+
past_key_values: Cache | None = None,
|
| 986 |
+
inputs_embeds: Tensor | None = None,
|
| 987 |
+
labels: Tensor | None = None,
|
| 988 |
+
use_cache: bool | None = None,
|
| 989 |
+
logits_to_keep: int | Tensor = 0,
|
| 990 |
+
output_attentions: bool | None = None,
|
| 991 |
+
output_hidden_states: bool | None = None,
|
| 992 |
+
return_dict: bool | None = None,
|
| 993 |
+
**kwargs: Any,
|
| 994 |
+
) -> CausalLMOutputWithPast | tuple[Tensor, ...]:
|
| 995 |
+
outputs = self.model(
|
| 996 |
+
input_ids=input_ids,
|
| 997 |
+
attention_mask=attention_mask,
|
| 998 |
+
position_ids=position_ids,
|
| 999 |
+
past_key_values=past_key_values,
|
| 1000 |
+
use_cache=use_cache,
|
| 1001 |
+
inputs_embeds=inputs_embeds,
|
| 1002 |
+
output_attentions=output_attentions,
|
| 1003 |
+
output_hidden_states=output_hidden_states,
|
| 1004 |
+
return_dict=True,
|
| 1005 |
+
**kwargs,
|
| 1006 |
+
)
|
| 1007 |
+
hidden_states = outputs.last_hidden_state
|
| 1008 |
+
indices = (
|
| 1009 |
+
slice(-logits_to_keep, None)
|
| 1010 |
+
if isinstance(logits_to_keep, int) and logits_to_keep > 0
|
| 1011 |
+
else logits_to_keep
|
| 1012 |
+
if isinstance(logits_to_keep, Tensor)
|
| 1013 |
+
else slice(None)
|
| 1014 |
+
)
|
| 1015 |
+
logits = self._softcapped_logits(hidden_states[:, indices, :])
|
| 1016 |
+
loss = None
|
| 1017 |
+
if labels is not None:
|
| 1018 |
+
shift_logits = logits[:, :-1].contiguous().float()
|
| 1019 |
+
shift_labels = labels[:, 1:].contiguous().to(shift_logits.device)
|
| 1020 |
+
loss = F.cross_entropy(
|
| 1021 |
+
shift_logits.view(-1, shift_logits.size(-1)),
|
| 1022 |
+
shift_labels.view(-1),
|
| 1023 |
+
ignore_index=-100,
|
| 1024 |
+
)
|
| 1025 |
+
if return_dict is False:
|
| 1026 |
+
result = (logits, outputs.past_key_values)
|
| 1027 |
+
return ((loss,) + result) if loss is not None else result
|
| 1028 |
+
return CausalLMOutputWithPast(
|
| 1029 |
+
loss=loss,
|
| 1030 |
+
logits=logits,
|
| 1031 |
+
past_key_values=outputs.past_key_values,
|
| 1032 |
+
hidden_states=outputs.hidden_states,
|
| 1033 |
+
attentions=outputs.attentions,
|
| 1034 |
+
)
|
| 1035 |
+
|
| 1036 |
+
|
| 1037 |
+
__all__ = [
|
| 1038 |
+
"LimiteDecoderLayer",
|
| 1039 |
+
"LimiteForCausalLM",
|
| 1040 |
+
"LimiteModel",
|
| 1041 |
+
"LimitePreTrainedModel",
|
| 1042 |
+
]
|
registration.py
ADDED
|
@@ -0,0 +1,77 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
"""Checkpoint-layout registration for the Hub-hosted Limite implementation."""
|
| 2 |
+
|
| 3 |
+
import torch
|
| 4 |
+
from transformers import Chunk, Concatenate, WeightConverter
|
| 5 |
+
from transformers.conversion_mapping import register_checkpoint_conversion_mapping
|
| 6 |
+
|
| 7 |
+
from .configuration_limite import LimiteConfig
|
| 8 |
+
|
| 9 |
+
|
| 10 |
+
class _SplitGQAQKV(Chunk):
|
| 11 |
+
"""Reverse a fused GQA projection without assuming equal Q/K/V sizes."""
|
| 12 |
+
|
| 13 |
+
@torch.no_grad()
|
| 14 |
+
def convert(
|
| 15 |
+
self,
|
| 16 |
+
input_dict: dict[str, torch.Tensor | list[torch.Tensor]],
|
| 17 |
+
source_patterns: list[str],
|
| 18 |
+
target_patterns: list[str],
|
| 19 |
+
*,
|
| 20 |
+
config: LimiteConfig,
|
| 21 |
+
**kwargs: object,
|
| 22 |
+
) -> dict[str, torch.Tensor]:
|
| 23 |
+
del source_patterns, kwargs
|
| 24 |
+
value = next(iter(input_dict.values()))
|
| 25 |
+
tensor = value[0] if isinstance(value, list) else value
|
| 26 |
+
q_size = int(config.num_attention_heads) * int(config.head_dim)
|
| 27 |
+
kv_size = int(config.num_key_value_heads) * int(config.head_dim)
|
| 28 |
+
expected = q_size + 2 * kv_size
|
| 29 |
+
if tensor.shape[self.dim] != expected:
|
| 30 |
+
raise ValueError(
|
| 31 |
+
"Invalid fused QKV size: "
|
| 32 |
+
f"expected {expected}, found {tensor.shape[self.dim]}."
|
| 33 |
+
)
|
| 34 |
+
chunks = tensor.split((q_size, kv_size, kv_size), dim=self.dim)
|
| 35 |
+
return dict(zip(target_patterns, chunks, strict=True))
|
| 36 |
+
|
| 37 |
+
@property
|
| 38 |
+
def reverse_op(self) -> "_ConcatenateGQAQKV":
|
| 39 |
+
return _ConcatenateGQAQKV(self.dim)
|
| 40 |
+
|
| 41 |
+
|
| 42 |
+
class _ConcatenateGQAQKV(Concatenate):
|
| 43 |
+
"""Fuse Q/K/V while retaining a GQA-aware reverse save transform."""
|
| 44 |
+
|
| 45 |
+
@property
|
| 46 |
+
def reverse_op(self) -> _SplitGQAQKV:
|
| 47 |
+
return _SplitGQAQKV(self.dim)
|
| 48 |
+
|
| 49 |
+
|
| 50 |
+
def register_weight_converters() -> None:
|
| 51 |
+
"""Register split-checkpoint to fused-runtime transformations."""
|
| 52 |
+
register_checkpoint_conversion_mapping(
|
| 53 |
+
LimiteConfig.model_type,
|
| 54 |
+
[
|
| 55 |
+
WeightConverter(
|
| 56 |
+
source_patterns=[
|
| 57 |
+
"self_attn.q_proj.weight",
|
| 58 |
+
"self_attn.k_proj.weight",
|
| 59 |
+
"self_attn.v_proj.weight",
|
| 60 |
+
],
|
| 61 |
+
target_patterns="self_attn.qkv_proj.weight",
|
| 62 |
+
operations=[_ConcatenateGQAQKV(dim=0)],
|
| 63 |
+
),
|
| 64 |
+
WeightConverter(
|
| 65 |
+
source_patterns=[
|
| 66 |
+
"mlp.gate_proj.weight",
|
| 67 |
+
"mlp.up_proj.weight",
|
| 68 |
+
],
|
| 69 |
+
target_patterns="mlp.gate_up_proj.weight",
|
| 70 |
+
operations=[Concatenate(dim=0)],
|
| 71 |
+
),
|
| 72 |
+
],
|
| 73 |
+
overwrite=True,
|
| 74 |
+
)
|
| 75 |
+
|
| 76 |
+
|
| 77 |
+
__all__ = ["register_weight_converters"]
|
tokenizer_config.json
CHANGED
|
@@ -215,6 +215,7 @@
|
|
| 215 |
"eos_token": "<|endoftext|>",
|
| 216 |
"errors": "replace",
|
| 217 |
"extra_special_tokens": {},
|
|
|
|
| 218 |
"model_max_length": 131072,
|
| 219 |
"pad_token": "<|endoftext|>",
|
| 220 |
"padding_side": "right",
|
|
|
|
| 215 |
"eos_token": "<|endoftext|>",
|
| 216 |
"errors": "replace",
|
| 217 |
"extra_special_tokens": {},
|
| 218 |
+
"fix_mistral_regex": false,
|
| 219 |
"model_max_length": 131072,
|
| 220 |
"pad_token": "<|endoftext|>",
|
| 221 |
"padding_side": "right",
|