Rishik001 commited on
Commit
4ec06ef
·
verified ·
1 Parent(s): b12a8fa

Upload step_003053 from MoE-bucket/checkpoints_kda_run_1308_12h_1b_6b

Browse files
config.json ADDED
@@ -0,0 +1,87 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "architectures": [
3
+ "LlamaKDA"
4
+ ],
5
+ "attention_bias": false,
6
+ "attention_dropout": 0.0,
7
+ "auto_map": {
8
+ "AutoConfig": "model.LlamaKDAConfig",
9
+ "AutoModelForCausalLM": "model.LlamaKDA"
10
+ },
11
+ "bos_token_id": 1,
12
+ "dtype": "bfloat16",
13
+ "eos_token_id": 2,
14
+ "head_dim": 128,
15
+ "hidden_act": "silu",
16
+ "hidden_size": 1536,
17
+ "initializer_range": 0.02,
18
+ "intermediate_size": 5120,
19
+ "kda_allow_neg_eigval": false,
20
+ "kda_conv_bias": false,
21
+ "kda_conv_size": 4,
22
+ "kda_every": 4,
23
+ "kda_expand_v": 1.0,
24
+ "kda_full_attn_every": null,
25
+ "kda_full_attn_layers": null,
26
+ "kda_full_attn_range": null,
27
+ "kda_head_dim": 128,
28
+ "kda_layer_types": [
29
+ "kda",
30
+ "full",
31
+ "full",
32
+ "full",
33
+ "kda",
34
+ "full",
35
+ "full",
36
+ "full",
37
+ "kda",
38
+ "full",
39
+ "full",
40
+ "full",
41
+ "kda",
42
+ "full",
43
+ "full",
44
+ "full",
45
+ "kda",
46
+ "full",
47
+ "full",
48
+ "full",
49
+ "kda",
50
+ "full",
51
+ "full",
52
+ "full",
53
+ "kda",
54
+ "full",
55
+ "full",
56
+ "full",
57
+ "kda",
58
+ "full",
59
+ "full",
60
+ "full"
61
+ ],
62
+ "kda_lower_bound": null,
63
+ "kda_num_heads": 12,
64
+ "kda_num_v_heads": null,
65
+ "kda_offset": 0,
66
+ "kda_safe_gate": false,
67
+ "kda_use_short_conv": true,
68
+ "max_position_embeddings": 8192,
69
+ "mlp_bias": false,
70
+ "model_type": "llama_kda",
71
+ "num_attention_heads": 12,
72
+ "num_hidden_layers": 32,
73
+ "num_key_value_heads": 6,
74
+ "pad_token_id": 2,
75
+ "pretraining_tp": 1,
76
+ "qk_norm": true,
77
+ "qk_norm_eps": 1e-06,
78
+ "rms_norm_eps": 1e-06,
79
+ "rope_parameters": {
80
+ "rope_theta": 1000000,
81
+ "rope_type": "default"
82
+ },
83
+ "tie_word_embeddings": true,
84
+ "transformers_version": "5.8.0",
85
+ "use_cache": false,
86
+ "vocab_size": 32000
87
+ }
generation_config.json ADDED
@@ -0,0 +1,10 @@
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "_from_model_config": true,
3
+ "bos_token_id": 1,
4
+ "eos_token_id": 2,
5
+ "output_attentions": false,
6
+ "output_hidden_states": false,
7
+ "pad_token_id": 2,
8
+ "transformers_version": "5.8.0",
9
+ "use_cache": false
10
+ }
model.py ADDED
@@ -0,0 +1,584 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ """Model architectures shared by the baseline and LongCat training scripts."""
2
+
3
+ import math
4
+ import warnings
5
+ from typing import Optional
6
+
7
+ import torch
8
+ import torch.nn as nn
9
+ from transformers import LlamaConfig, LlamaForCausalLM
10
+ from transformers.models.llama import modeling_llama as llama_modeling
11
+
12
+ try:
13
+ # Kimi Delta Attention layer (linear attention) from flash-linear-attention.
14
+ # Optional: only the KDA architecture needs it, so the baseline/longcat scripts
15
+ # keep importing model.py even when FLA is absent.
16
+ from fla.layers.kda import KimiDeltaAttention
17
+
18
+ FLA_KDA_IMPORT_ERROR = None
19
+ except ImportError as exc: # pragma: no cover - exercised only without FLA
20
+ KimiDeltaAttention = None
21
+ FLA_KDA_IMPORT_ERROR = exc
22
+
23
+
24
+ # Distinct prime hash multipliers ("salts"), one per (order, head) table. Giving
25
+ # every head a different polynomial base makes the K hash functions genuinely
26
+ # independent *regardless of table size*, so the historical K->1 collapse — two
27
+ # heads that shared a table size hashed every n-gram to the identical slot —
28
+ # cannot recur. Primes larger than the base vocab keep each per-head polynomial
29
+ # injective over the token range, and being coprime to the table sizes avoids the
30
+ # base-multiple collision spike the LongCat paper reports (Fig. 3b): the effective
31
+ # base (multiplier mod table_size) is then a scrambled value rather than the raw
32
+ # vocab size. All of this is enforced at build time by _validate_ngram_hashing.
33
+ _DEFAULT_HASH_MULTIPLIERS = (
34
+ 40009, 100003, 262147, 524287, 1000003, 2000003,
35
+ 3000017, 4000037, 5000011, 6000101, 7000127, 8000009,
36
+ )
37
+
38
+ # Default minimum pairwise separation between table sizes (fraction of the
39
+ # smaller size). Near-equal sizes are the condition that silently disabled
40
+ # multi-head hashing before, so the build refuses to start below this.
41
+ _DEFAULT_MIN_PAIRWISE_SIZE_GAP = 0.005
42
+
43
+
44
+ def _validate_ngram_hashing(
45
+ table_sizes: list[int],
46
+ multipliers: list[int],
47
+ base_vocab: int,
48
+ min_pairwise_size_gap: float,
49
+ ) -> None:
50
+ """Refuse to build a degenerate n-gram hashing setup. Raises ValueError.
51
+
52
+ Guards, in order of how badly they corrupt the experiment:
53
+
54
+ 1. No two tables may be the *same hash function*. Two heads are identical iff
55
+ they share both a table size and an effective base (multiplier mod size);
56
+ that is the exact K->1 collapse. This is the load-bearing invariant.
57
+ 2. Each multiplier must be >= base vocab, or distinct n-grams alias before the
58
+ modulus (the base-`m` polynomial stops being injective over token digits).
59
+ 3. Each multiplier must be coprime to its table size, or one n-gram coordinate
60
+ collapses into gcd-many classes (the mechanism behind the paper's spike).
61
+ 4. Table sizes must not be near-duplicates — the config that hid the clone bug.
62
+
63
+ A soft warning also fires when a size sits within 5% of base vocab of an
64
+ integer multiple of it (paper Fig. 3b), which prime multipliers mitigate but
65
+ do not fully erase.
66
+ """
67
+ n = len(table_sizes)
68
+ if len(multipliers) != n:
69
+ raise ValueError(
70
+ f"Expected {n} ngram hash multipliers (one per table), got {len(multipliers)}."
71
+ )
72
+
73
+ # (1) precise clone check: identical (size, effective base) => identical indices.
74
+ seen: dict[tuple[int, int], int] = {}
75
+ for idx, (m, s) in enumerate(zip(multipliers, table_sizes)):
76
+ key = (s, m % s)
77
+ if key in seen:
78
+ raise ValueError(
79
+ f"N-gram hash tables {seen[key]} and {idx} are the SAME hash function "
80
+ f"(table size {s}, effective base {m % s}). Two heads that hash "
81
+ "identically collapse K sub-tables to K=1 — exactly the bug this guard "
82
+ "exists to prevent. Give them distinct multipliers or distinct sizes."
83
+ )
84
+ seen[key] = idx
85
+
86
+ for m, s in zip(multipliers, table_sizes):
87
+ # (2) injectivity of the base-`m` polynomial over token digits [0, base_vocab).
88
+ if m < base_vocab:
89
+ raise ValueError(
90
+ f"N-gram hash multiplier {m} must be >= base vocab {base_vocab}; "
91
+ "a smaller base aliases distinct n-grams before the modulus is applied."
92
+ )
93
+ # (3) coprimality: gcd > 1 collapses a coordinate into gcd-many residues.
94
+ g = math.gcd(m, s)
95
+ if g != 1:
96
+ raise ValueError(
97
+ f"N-gram hash multiplier {m} shares factor {g} with table size {s}. "
98
+ "Pick a multiplier coprime to the table size (a prime larger than every "
99
+ "table size is always safe) so no n-gram coordinate collapses."
100
+ )
101
+
102
+ # (4) near-duplicate sizes: the historical trigger for the clone collapse.
103
+ order = sorted(range(n), key=lambda i: table_sizes[i])
104
+ for a, b in zip(order, order[1:]):
105
+ sa, sb = table_sizes[a], table_sizes[b]
106
+ rel = abs(sa - sb) / min(sa, sb)
107
+ if rel < min_pairwise_size_gap:
108
+ raise ValueError(
109
+ f"N-gram table sizes {sa} and {sb} differ by only {rel * 100:.3f}% "
110
+ f"(guard requires >= {min_pairwise_size_gap * 100:.3f}%). Near-equal "
111
+ "sizes are the condition that silently disabled multi-head hashing "
112
+ "before; spread the table sizes apart."
113
+ )
114
+
115
+ # (5) soft: sizes near an integer multiple of base vocab (paper Fig. 3b).
116
+ for s in table_sizes:
117
+ dist = min(s % base_vocab, base_vocab - s % base_vocab)
118
+ if dist / base_vocab < 0.05:
119
+ warnings.warn(
120
+ f"N-gram table size {s} is within {dist} of an integer multiple of base "
121
+ f"vocab {base_vocab}; the LongCat paper reports collision spikes there. "
122
+ "The prime multipliers mitigate this, but consider nudging the size.",
123
+ stacklevel=2,
124
+ )
125
+
126
+
127
+ class LlamaLongCatNgramConfig(LlamaConfig):
128
+ """Serializable configuration for :class:`LlamaLongCatNgram`."""
129
+
130
+ model_type = "llama_longcat_ngram"
131
+
132
+ def __init__(
133
+ self,
134
+ ngram_max_n: int = 4,
135
+ ngram_num_heads: int = 2,
136
+ ngram_table_vocab_sizes: Optional[list[int]] = None,
137
+ ngram_embedding_amplification: str = "layer_norm",
138
+ ngram_hash_multipliers: Optional[list[int]] = None,
139
+ ngram_min_pairwise_size_gap: float = _DEFAULT_MIN_PAIRWISE_SIZE_GAP,
140
+ qk_norm: bool = False,
141
+ qk_norm_eps: Optional[float] = None,
142
+ **kwargs,
143
+ ):
144
+ super().__init__(**kwargs)
145
+ self.ngram_max_n = ngram_max_n
146
+ self.ngram_num_heads = ngram_num_heads
147
+ self.ngram_table_vocab_sizes = ngram_table_vocab_sizes
148
+ self.ngram_embedding_amplification = ngram_embedding_amplification
149
+ self.ngram_hash_multipliers = ngram_hash_multipliers
150
+ self.ngram_min_pairwise_size_gap = ngram_min_pairwise_size_gap
151
+ self.qk_norm = qk_norm
152
+ self.qk_norm_eps = qk_norm_eps
153
+
154
+
155
+ class LongCatNgramEmbedder(nn.Module):
156
+ """LongCat N-gram Embedding from Eq. 2 and Eq. 3 of arXiv:2601.21204."""
157
+
158
+ def __init__(self, config: LlamaLongCatNgramConfig):
159
+ super().__init__()
160
+ self.max_n = config.ngram_max_n
161
+ self.num_heads = config.ngram_num_heads
162
+ self.base_vocab_size = config.vocab_size
163
+ self.eos_token_id = config.eos_token_id
164
+ self.orders = list(range(2, self.max_n + 1))
165
+
166
+ num_tables = len(self.orders) * self.num_heads
167
+ table_vocab_sizes = config.ngram_table_vocab_sizes
168
+ if config.hidden_size % num_tables != 0:
169
+ raise ValueError(
170
+ f"hidden_size ({config.hidden_size}) must be divisible by "
171
+ f"(ngram_max_n-1)*ngram_num_heads ({num_tables})."
172
+ )
173
+ if table_vocab_sizes is None or len(table_vocab_sizes) != num_tables:
174
+ actual = None if table_vocab_sizes is None else len(table_vocab_sizes)
175
+ raise ValueError(
176
+ f"Expected {num_tables} ngram_table_vocab_sizes "
177
+ f"((ngram_max_n-1)*ngram_num_heads), got {actual}."
178
+ )
179
+ self.sub_dim = config.hidden_size // num_tables
180
+
181
+ # Resolve the per-head hash multipliers (built-in defaults unless the
182
+ # config overrides them) and refuse to build a degenerate setup.
183
+ configured = config.ngram_hash_multipliers
184
+ if configured is None:
185
+ if num_tables > len(_DEFAULT_HASH_MULTIPLIERS):
186
+ raise ValueError(
187
+ f"Need {num_tables} hash multipliers but only "
188
+ f"{len(_DEFAULT_HASH_MULTIPLIERS)} defaults are defined; pass "
189
+ "ngram_hash_multipliers explicitly."
190
+ )
191
+ configured = _DEFAULT_HASH_MULTIPLIERS[:num_tables]
192
+ multipliers: list[int] = [int(m) for m in configured]
193
+ _validate_ngram_hashing(
194
+ list(table_vocab_sizes),
195
+ multipliers,
196
+ self.base_vocab_size,
197
+ config.ngram_min_pairwise_size_gap,
198
+ )
199
+ # Persist the resolved list so it is serialized in config.json.
200
+ config.ngram_hash_multipliers = multipliers
201
+
202
+ self.tables = nn.ModuleDict()
203
+ self.projections = nn.ModuleDict()
204
+ self.multipliers: dict[str, int] = {}
205
+ idx = 0
206
+ for n in self.orders:
207
+ for k in range(self.num_heads):
208
+ key = f"n{n}_k{k}"
209
+ self.tables[key] = nn.Embedding(
210
+ table_vocab_sizes[idx], self.sub_dim
211
+ )
212
+ self.projections[key] = nn.Linear(
213
+ self.sub_dim, config.hidden_size, bias=False
214
+ )
215
+ self.multipliers[key] = multipliers[idx]
216
+ idx += 1
217
+
218
+ amplification = config.ngram_embedding_amplification.strip().lower()
219
+ if amplification == "layer_norm":
220
+ self.amplification = nn.LayerNorm(config.hidden_size)
221
+ self.amplification_scale = 1.0
222
+ elif amplification == "sqrt_d":
223
+ self.amplification = nn.Identity()
224
+ self.amplification_scale = math.sqrt(config.hidden_size)
225
+ elif amplification == "none":
226
+ self.amplification = nn.Identity()
227
+ self.amplification_scale = 1.0
228
+ else:
229
+ raise ValueError(
230
+ "ngram_embedding_amplification must be one of "
231
+ "{'layer_norm', 'sqrt_d', 'none'}, got "
232
+ f"{amplification!r}."
233
+ )
234
+
235
+ def _shift_right(self, x: torch.Tensor, shift: int) -> torch.Tensor:
236
+ """Causal shift, zeroing context that crosses an EOS boundary."""
237
+ if shift == 0:
238
+ return x
239
+ pad = x.new_zeros(x.shape[0], shift)
240
+ shifted = torch.cat([pad, x[:, :-shift]], dim=1)
241
+
242
+ crosses_eos = torch.zeros_like(x, dtype=torch.bool)
243
+ for offset in range(1, shift + 1):
244
+ previous = torch.cat(
245
+ [x.new_zeros(x.shape[0], offset), x[:, :-offset]], dim=1
246
+ )
247
+ crosses_eos |= previous.eq(self.eos_token_id)
248
+ return shifted.masked_fill(crosses_eos, 0)
249
+
250
+ def _hash_ngram(
251
+ self,
252
+ input_ids: torch.Tensor,
253
+ n: int,
254
+ table_size: int,
255
+ multiplier: int,
256
+ shifted_tokens: Optional[dict[int, torch.Tensor]] = None,
257
+ ) -> torch.Tensor:
258
+ """Eq. 2 with a per-head base: sum_j t[i-j] * multiplier**j mod table_size.
259
+
260
+ `multiplier` is this head's hash salt (a distinct prime >= base vocab), so
261
+ two heads never compute the same indices even at equal table sizes. The
262
+ modulus is applied every Horner step, so the result matches the full
263
+ polynomial mod `table_size` while staying far inside int64.
264
+ """
265
+ h = torch.zeros_like(input_ids)
266
+ for j in range(n - 1, -1, -1):
267
+ tok = (
268
+ input_ids
269
+ if j == 0
270
+ else shifted_tokens[j]
271
+ if shifted_tokens is not None
272
+ else self._shift_right(input_ids, j)
273
+ )
274
+ h = (h * multiplier + tok) % table_size
275
+ return h
276
+
277
+ def forward(
278
+ self, input_ids: torch.Tensor, base_embeddings: torch.Tensor
279
+ ) -> torch.Tensor:
280
+ """Return amplified Eq. 3 embeddings with shape [B, T, H]."""
281
+ combined = base_embeddings
282
+ shifted_tokens = {
283
+ shift: self._shift_right(input_ids, shift)
284
+ for shift in range(1, self.max_n)
285
+ }
286
+ for n in self.orders:
287
+ for k in range(self.num_heads):
288
+ key = f"n{n}_k{k}"
289
+ table_size = self.tables[key].num_embeddings
290
+ hash_ids = self._hash_ngram(
291
+ input_ids, n, table_size, self.multipliers[key], shifted_tokens
292
+ )
293
+ combined = combined + self.projections[key](
294
+ self.tables[key](hash_ids)
295
+ )
296
+
297
+ combined = combined / (len(self.orders) * self.num_heads + 1)
298
+ return self.amplification(combined) * self.amplification_scale
299
+
300
+
301
+ class LlamaQKNormAttention(llama_modeling.LlamaAttention):
302
+ """LLaMA attention with per-head RMSNorm on Q and K before RoPE."""
303
+
304
+ def __init__(self, config: LlamaLongCatNgramConfig, layer_idx: int):
305
+ super().__init__(config, layer_idx)
306
+ eps = config.qk_norm_eps if config.qk_norm_eps is not None else config.rms_norm_eps
307
+ self.q_norm = llama_modeling.LlamaRMSNorm(self.head_dim, eps=eps)
308
+ self.k_norm = llama_modeling.LlamaRMSNorm(self.head_dim, eps=eps)
309
+
310
+ def forward(
311
+ self, hidden_states: torch.Tensor, position_embeddings=None,
312
+ attention_mask=None, past_key_values=None, **kwargs,
313
+ ):
314
+ input_shape = hidden_states.shape[:-1]
315
+ hidden_shape = (*input_shape, -1, self.head_dim)
316
+ query_states = self.q_norm(self.q_proj(hidden_states).view(hidden_shape)).transpose(1, 2)
317
+ key_states = self.k_norm(self.k_proj(hidden_states).view(hidden_shape)).transpose(1, 2)
318
+ value_states = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2)
319
+
320
+ cos, sin = position_embeddings
321
+ query_states, key_states = llama_modeling.apply_rotary_pos_emb(
322
+ query_states, key_states, cos, sin
323
+ )
324
+ if past_key_values is not None:
325
+ key_states, value_states = past_key_values.update(
326
+ key_states, value_states, self.layer_idx
327
+ )
328
+
329
+ attention_interface = llama_modeling.ALL_ATTENTION_FUNCTIONS.get_interface(
330
+ self.config._attn_implementation, llama_modeling.eager_attention_forward
331
+ )
332
+ attn_output, attn_weights = attention_interface(
333
+ self, query_states, key_states, value_states, attention_mask,
334
+ dropout=0.0 if not self.training else self.attention_dropout,
335
+ scaling=self.scaling, **kwargs,
336
+ )
337
+ attn_output = self.o_proj(attn_output.reshape(*input_shape, -1).contiguous())
338
+ return attn_output, attn_weights
339
+
340
+
341
+ class LlamaLongCatNgram(LlamaForCausalLM):
342
+ """LLaMA using LongCat's standard input N-gram Embedding (NE)."""
343
+
344
+ config_class = LlamaLongCatNgramConfig
345
+
346
+ def __init__(self, config: LlamaLongCatNgramConfig):
347
+ super().__init__(config)
348
+ if config.qk_norm:
349
+ for layer_idx, layer in enumerate(self.model.layers):
350
+ layer.self_attn = LlamaQKNormAttention(config, layer_idx)
351
+ self.ngram_embedder = LongCatNgramEmbedder(config)
352
+
353
+ def forward(self, input_ids=None, inputs_embeds=None, **kwargs):
354
+ if input_ids is not None and inputs_embeds is not None:
355
+ raise ValueError("Specify exactly one of input_ids or inputs_embeds.")
356
+ if input_ids is not None:
357
+ base_embeddings = self.model.embed_tokens(input_ids)
358
+ inputs_embeds = self.ngram_embedder(input_ids, base_embeddings)
359
+ input_ids = None
360
+ return super().forward(
361
+ input_ids=input_ids, inputs_embeds=inputs_embeds, **kwargs
362
+ )
363
+
364
+ def prepare_inputs_for_generation(self, input_ids, **kwargs):
365
+ """Preserve causal n-gram context when HF slices cached decode inputs."""
366
+ model_inputs = super().prepare_inputs_for_generation(input_ids, **kwargs)
367
+ prepared_ids = model_inputs.get("input_ids")
368
+ if prepared_ids is None:
369
+ return model_inputs
370
+
371
+ base_embeddings = self.model.embed_tokens(input_ids)
372
+ full_embeddings = self.ngram_embedder(input_ids, base_embeddings)
373
+ model_inputs["inputs_embeds"] = full_embeddings[:, -prepared_ids.shape[1] :]
374
+ model_inputs["input_ids"] = None
375
+ return model_inputs
376
+
377
+
378
+ # Instruct save_pretrained() to package this source file and write AutoClass
379
+ # metadata. Loading the resulting checkpoint requires trust_remote_code=True.
380
+ LlamaLongCatNgramConfig.register_for_auto_class()
381
+ LlamaLongCatNgram.register_for_auto_class("AutoModelForCausalLM")
382
+
383
+
384
+ class LlamaKDAConfig(LlamaConfig):
385
+ """Config for :class:`LlamaKDA` — a hybrid Kimi-Delta / softmax LLaMA.
386
+
387
+ A subset of decoder layers use Kimi Delta Attention (KDA, a gated-delta
388
+ linear attention); the rest keep standard softmax self-attention (with the
389
+ same FA backend and optional QK-norm as the baseline). Which layers are which
390
+ is resolved by :func:`resolve_kda_layer_types`, honoring (in priority order)
391
+ ``kda_full_attn_layers`` > ``kda_full_attn_range`` > ``kda_full_attn_every``.
392
+ With none set, every layer is KDA (pure linear attention).
393
+ """
394
+
395
+ model_type = "llama_kda"
396
+
397
+ def __init__(
398
+ self,
399
+ # --- Hybrid layout: which layers keep softmax (full) attention ---
400
+ kda_full_attn_layers: Optional[list[int]] = None,
401
+ kda_full_attn_every: Optional[int] = None,
402
+ kda_full_attn_range: Optional[list[int]] = None,
403
+ # Sparse-KDA interleave (the inverse of kda_full_attn_every): one KDA layer
404
+ # every `kda_every` layers, at indices where i % kda_every == kda_offset;
405
+ # every other layer is full (GQA) attention.
406
+ kda_every: Optional[int] = None,
407
+ kda_offset: int = 0,
408
+ # --- KDA layer hyperparameters (forwarded to fla KimiDeltaAttention) ---
409
+ kda_head_dim: int = 128,
410
+ kda_num_heads: Optional[int] = None,
411
+ kda_num_v_heads: Optional[int] = None,
412
+ kda_expand_v: float = 1.0,
413
+ kda_use_short_conv: bool = True,
414
+ kda_conv_size: int = 4,
415
+ kda_conv_bias: bool = False,
416
+ kda_allow_neg_eigval: bool = False,
417
+ kda_lower_bound: Optional[float] = None,
418
+ kda_safe_gate: bool = False,
419
+ # --- QK-norm applies to the softmax (full-attention) layers only ---
420
+ qk_norm: bool = False,
421
+ qk_norm_eps: Optional[float] = None,
422
+ **kwargs,
423
+ ):
424
+ super().__init__(**kwargs)
425
+ self.kda_full_attn_layers = kda_full_attn_layers
426
+ self.kda_full_attn_every = kda_full_attn_every
427
+ self.kda_full_attn_range = kda_full_attn_range
428
+ self.kda_every = kda_every
429
+ self.kda_offset = kda_offset
430
+ self.kda_head_dim = kda_head_dim
431
+ self.kda_num_heads = kda_num_heads
432
+ self.kda_num_v_heads = kda_num_v_heads
433
+ self.kda_expand_v = kda_expand_v
434
+ self.kda_use_short_conv = kda_use_short_conv
435
+ self.kda_conv_size = kda_conv_size
436
+ self.kda_conv_bias = kda_conv_bias
437
+ self.kda_allow_neg_eigval = kda_allow_neg_eigval
438
+ self.kda_lower_bound = kda_lower_bound
439
+ self.kda_safe_gate = kda_safe_gate
440
+ self.qk_norm = qk_norm
441
+ self.qk_norm_eps = qk_norm_eps
442
+
443
+
444
+ def resolve_kda_layer_types(config: LlamaKDAConfig) -> list[str]:
445
+ """Return a per-layer list of ``"kda"`` / ``"full"`` (softmax) attention.
446
+
447
+ Priority: explicit ``kda_full_attn_layers`` > contiguous ``kda_full_attn_range``
448
+ ``[start, end)`` > interleaved ``kda_full_attn_every`` (the last layer of every
449
+ block of ``n`` is full attention, e.g. ``4`` -> Kimi/Qwen-style 3:1) > ``kda_every``
450
+ (the INVERSE — KDA is the sparse type: one KDA layer every ``kda_every`` layers at
451
+ ``i % kda_every == kda_offset``, all others full). If none is set, all layers are KDA.
452
+ """
453
+ n = config.num_hidden_layers
454
+ if config.kda_full_attn_layers is not None:
455
+ full = set(int(i) for i in config.kda_full_attn_layers)
456
+ elif config.kda_full_attn_range is not None:
457
+ start, end = config.kda_full_attn_range
458
+ full = set(range(int(start), int(end)))
459
+ elif config.kda_full_attn_every:
460
+ every = int(config.kda_full_attn_every)
461
+ if every < 1:
462
+ raise ValueError(f"kda_full_attn_every must be >= 1, got {every}.")
463
+ full = {i for i in range(n) if (i + 1) % every == 0}
464
+ elif getattr(config, "kda_every", None):
465
+ # Inverse of kda_full_attn_every: KDA is the SPARSE type. One KDA layer every
466
+ # `kda_every` layers at i % kda_every == kda_offset; every other layer is full.
467
+ every = int(config.kda_every)
468
+ if every < 1:
469
+ raise ValueError(f"kda_every must be >= 1, got {every}.")
470
+ offset = int(getattr(config, "kda_offset", 0) or 0) % every
471
+ kda = {i for i in range(n) if i % every == offset}
472
+ full = set(range(n)) - kda
473
+ else:
474
+ full = set()
475
+ for i in full:
476
+ if not 0 <= i < n:
477
+ raise ValueError(
478
+ f"Full-attention layer index {i} is out of range for "
479
+ f"num_hidden_layers={n}."
480
+ )
481
+ return ["full" if i in full else "kda" for i in range(n)]
482
+
483
+
484
+ class LlamaKDAAttention(nn.Module):
485
+ """Adapter wrapping fla's :class:`KimiDeltaAttention` for a LLaMA decoder layer.
486
+
487
+ KDA is linear attention: it carries no RoPE and normalizes q/k internally
488
+ (L2-norm), so ``position_embeddings`` are ignored here. The decoder layer
489
+ expects a ``(hidden_states, attn_weights)`` pair back; KDA returns a triple, so
490
+ we drop the cache/weights. A 4-D causal mask (built by ``LlamaModel`` for the
491
+ softmax layers) is meaningless to KDA — only a 2-D ``[B, T]`` padding mask is
492
+ forwarded; anything else becomes ``None`` (packed training carries no padding).
493
+
494
+ The module is run **stateless**: it never reads or writes ``past_key_values``.
495
+ HF's ``LlamaModel`` hands every layer an HF ``DynamicCache`` (incompatible with
496
+ fla's recurrent-state cache), which is fine for full-sequence LM-loss training
497
+ and eval but means this wrapper does not support HF incremental ``generate``.
498
+ """
499
+
500
+ def __init__(self, config: LlamaKDAConfig, layer_idx: int):
501
+ super().__init__()
502
+ if KimiDeltaAttention is None:
503
+ raise ImportError(
504
+ "LlamaKDA requires flash-linear-attention (fla) for KimiDeltaAttention."
505
+ ) from FLA_KDA_IMPORT_ERROR
506
+ head_dim = config.kda_head_dim
507
+ num_heads = config.kda_num_heads or (config.hidden_size // head_dim)
508
+ if num_heads * head_dim != config.hidden_size:
509
+ # fla supports q/k dim != hidden; Kimi-Linear over-provisions ~1.8x.
510
+ warnings.warn(
511
+ f"KDA q/k dim {num_heads * head_dim} != hidden_size "
512
+ f"{config.hidden_size}; layer params/state will differ from a "
513
+ "same-width softmax layer.",
514
+ stacklevel=2,
515
+ )
516
+ self.layer_idx = layer_idx
517
+ self.kda = KimiDeltaAttention(
518
+ hidden_size=config.hidden_size,
519
+ expand_v=config.kda_expand_v,
520
+ head_dim=head_dim,
521
+ num_heads=num_heads,
522
+ num_v_heads=config.kda_num_v_heads,
523
+ mode="chunk",
524
+ use_short_conv=config.kda_use_short_conv,
525
+ conv_size=config.kda_conv_size,
526
+ conv_bias=config.kda_conv_bias,
527
+ allow_neg_eigval=config.kda_allow_neg_eigval,
528
+ safe_gate=config.kda_safe_gate,
529
+ lower_bound=config.kda_lower_bound,
530
+ layer_idx=layer_idx,
531
+ norm_eps=config.rms_norm_eps,
532
+ )
533
+
534
+ def forward(
535
+ self, hidden_states: torch.Tensor, position_embeddings=None,
536
+ attention_mask=None, past_key_values=None, use_cache=False, **kwargs,
537
+ ):
538
+ mask = attention_mask if (attention_mask is not None and attention_mask.dim() == 2) else None
539
+ forward_kwargs = {k: v for k, v in kwargs.items() if k == "cu_seqlens"}
540
+ attn_output, _, _ = self.kda(
541
+ hidden_states=hidden_states,
542
+ attention_mask=mask,
543
+ past_key_values=None,
544
+ use_cache=False,
545
+ **forward_kwargs,
546
+ )
547
+ return attn_output, None
548
+
549
+
550
+ class LlamaKDA(LlamaForCausalLM):
551
+ """LLaMA whose attention is a config-driven hybrid of KDA and softmax layers."""
552
+
553
+ config_class = LlamaKDAConfig
554
+
555
+ def __init__(self, config: LlamaKDAConfig):
556
+ super().__init__(config)
557
+ layer_types = resolve_kda_layer_types(config)
558
+ for layer_idx, layer in enumerate(self.model.layers):
559
+ if layer_types[layer_idx] == "kda":
560
+ layer.self_attn = LlamaKDAAttention(config, layer_idx)
561
+ elif config.qk_norm:
562
+ layer.self_attn = LlamaQKNormAttention(config, layer_idx)
563
+ # otherwise keep the default softmax LlamaAttention from super().__init__.
564
+ # Persist the resolved layout so it lands in config.json and can be logged.
565
+ config.kda_layer_types = layer_types
566
+
567
+
568
+ LlamaKDAConfig.register_for_auto_class()
569
+ LlamaKDA.register_for_auto_class("AutoModelForCausalLM")
570
+
571
+
572
+ __all__ = [
573
+ "LlamaConfig",
574
+ "LlamaForCausalLM",
575
+ "LlamaLongCatNgramConfig",
576
+ "LongCatNgramEmbedder",
577
+ "LlamaQKNormAttention",
578
+ "LlamaLongCatNgram",
579
+ "LlamaKDAConfig",
580
+ "LlamaKDAAttention",
581
+ "LlamaKDA",
582
+ "resolve_kda_layer_types",
583
+ "_validate_ngram_hashing",
584
+ ]
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:505a143238c1c13889ca27d5bb2802c11665b9aff667a8f7877140ee98a83cf8
3
+ size 2112497840
tokenizer.json ADDED
The diff for this file is too large to render. See raw diff
 
tokenizer_config.json ADDED
@@ -0,0 +1,15 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "backend": "tokenizers",
3
+ "bos_token": "<s>",
4
+ "eos_token": "</s>",
5
+ "extra_special_tokens": [
6
+ "<s>",
7
+ "</s>"
8
+ ],
9
+ "is_local": true,
10
+ "local_files_only": false,
11
+ "model_max_length": 8192,
12
+ "pad_token": "</s>",
13
+ "tokenizer_class": "TokenizersBackend",
14
+ "unk_token": "<unk>"
15
+ }
training_config.json ADDED
@@ -0,0 +1,96 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "allow_inexact_legacy_data_resume": false,
3
+ "attention_bias": false,
4
+ "attn_implementation": "flash_attention_4",
5
+ "aux_adam_beta1": 0.9,
6
+ "aux_adam_beta2": 0.95,
7
+ "aux_adam_eps": 1e-08,
8
+ "aux_adam_lr": 0.0003,
9
+ "aux_adam_weight_decay": 0.1,
10
+ "beta1": 0.9,
11
+ "beta2": 0.95,
12
+ "checkpoint_bucket_folder": "checkpoints_llama_1b_kda_3to1_h12_6b_1308",
13
+ "cross_document_attention": true,
14
+ "dataloader_prefetch_factor": 2,
15
+ "dataloader_workers": 8,
16
+ "dataset_config": "sample-10BT",
17
+ "dataset_name": "HuggingFaceFW/fineweb",
18
+ "eval_benchmarks": false,
19
+ "eval_every_steps": 250,
20
+ "eval_holdout_fraction": 0.005,
21
+ "eval_max_examples": 512,
22
+ "eval_streaming_buffer_size": 1,
23
+ "grad_accum_steps": 24,
24
+ "grad_clip": 1.0,
25
+ "hf_assets_dir": "./hf_assets_llama_1b_kda_3to1_h12",
26
+ "hidden_act": "silu",
27
+ "hidden_size": 1536,
28
+ "init_from_checkpoint": null,
29
+ "initializer_range": 0.02,
30
+ "intermediate_size": 5120,
31
+ "kda_allow_neg_eigval": false,
32
+ "kda_conv_bias": false,
33
+ "kda_conv_size": 4,
34
+ "kda_every": 4,
35
+ "kda_expand_v": 1.0,
36
+ "kda_full_attn_every": null,
37
+ "kda_full_attn_layers": null,
38
+ "kda_full_attn_range": null,
39
+ "kda_head_dim": 128,
40
+ "kda_lower_bound": null,
41
+ "kda_num_heads": 12,
42
+ "kda_num_v_heads": null,
43
+ "kda_offset": 0,
44
+ "kda_safe_gate": false,
45
+ "kda_use_short_conv": true,
46
+ "learning_rate": 0.0003,
47
+ "liger_kernel_config": {
48
+ "cross_entropy": false,
49
+ "fused_linear_cross_entropy": true,
50
+ "rms_norm": true,
51
+ "rope": true,
52
+ "swiglu": true
53
+ },
54
+ "log_every_steps": 1,
55
+ "lr_decay_ratio": null,
56
+ "lr_decay_type": "cosine",
57
+ "max_position_embeddings": 8192,
58
+ "max_seq_len": 8192,
59
+ "max_steps": 3053,
60
+ "min_lr_factor": 0.1,
61
+ "muon_lr": 0.02,
62
+ "muon_momentum": 0.95,
63
+ "muon_nesterov": true,
64
+ "muon_ns_steps": 5,
65
+ "muon_weight_decay": 0.1,
66
+ "num_attention_heads": 12,
67
+ "num_hidden_layers": 32,
68
+ "num_key_value_heads": 6,
69
+ "optimizer_eps": 1e-08,
70
+ "optimizer_implementation": "fused",
71
+ "optimizer_name": "muon",
72
+ "output_dir": "./checkpoints_llama_1b_kda_3to1_h12_6b",
73
+ "per_device_batch_size": 10,
74
+ "qk_norm": true,
75
+ "qk_norm_eps": 1e-06,
76
+ "resume_from_checkpoint": null,
77
+ "rms_norm_eps": 1e-06,
78
+ "rope_theta": 1000000,
79
+ "save_every_steps": 250,
80
+ "seed": 42,
81
+ "streaming_buffer_size": 10000,
82
+ "sync_checkpoints_to_bucket": false,
83
+ "target_train_tokens": null,
84
+ "tie_word_embeddings": true,
85
+ "tokenizer_name": "meta-llama/Llama-2-7b",
86
+ "torch_compile": false,
87
+ "torch_compile_mode": "default",
88
+ "use_liger_kernel": true,
89
+ "use_wandb": true,
90
+ "vocab_size": 32000,
91
+ "wandb_log_console": false,
92
+ "wandb_project": "llama-1b-6b-torchtitan",
93
+ "wandb_run_name": "llama-1b-kda-3to1-h12-1308",
94
+ "warmup_steps": 150,
95
+ "weight_decay": 0.1
96
+ }