MisterOss commited on
Commit
b60c6b3
·
verified ·
1 Parent(s): 9402e42

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

Files changed (7) hide show
  1. README.md +54 -1
  2. config.json +5 -0
  3. configuration_limite.py +505 -0
  4. contract.py +16 -0
  5. modeling_limite.py +1042 -0
  6. registration.py +77 -0
  7. 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
- **Hugging Face Transformers compatibility is coming soon.** We will release support for loading and running Violetto directly with Transformers.
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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",