ArushBuilds commited on
Commit
4b04e69
Β·
1 Parent(s): 507e1d8

Update model.py

Browse files
Files changed (1) hide show
  1. model.py +240 -25
model.py CHANGED
@@ -26,6 +26,7 @@ _USE_LIGER = bool(getattr(alpha_config, 'use_liger', False))
26
 
27
  _LigerRMSNormFn = None
28
  try:
 
29
  from liger_kernel.ops.rms_norm import LigerRMSNormFunction as _LigerRMSNormFn
30
  except ImportError:
31
  pass
@@ -53,6 +54,15 @@ except ImportError: # pragma: no cover - fallback for package-style imports
53
 
54
  _SIMPLE_BLOCK_PATH = not _USE_FLASH_OPS
55
 
 
 
 
 
 
 
 
 
 
56
 
57
 
58
  _alpha_config_warned: set[str] = set()
@@ -98,7 +108,7 @@ _NVFP4_DISABLED_REASON: "str | None" = None
98
 
99
  _NVFP4_SKIP_KEYWORDS: tuple[str, ...] = (
100
  "wte", "lm_head", "ln_", "ln_f", "norm",
101
- "router", "gate", "expert_bias",
102
  )
103
 
104
 
@@ -235,7 +245,9 @@ def apply_int8_to_model(model: nn.Module, group_size: int = 128) -> int:
235
  return True
236
 
237
 
238
- before_ids = {fqn: id(m) for fqn, m in model.named_modules() if _filter_with_gs(m, fqn)}
 
 
239
 
240
 
241
  if _INT8_AVAILABLE and _Int8Config is not None:
@@ -253,13 +265,16 @@ def apply_int8_to_model(model: nn.Module, group_size: int = 128) -> int:
253
  _torchao_quantize(model, config_obj, filter_fn=_filter_with_gs)
254
 
255
 
256
- count = sum(
257
- 1 for fqn, old_id in before_ids.items()
258
- if id(dict(model.named_modules()).get(fqn)) != old_id
259
- )
 
 
260
  model._quant_mode = "int8"
261
  model._use_int8 = True
262
  model._use_nvfp4 = False
 
263
 
264
  def _alpha_config_attr(name: str, default):
265
  if not hasattr(alpha_config, name):
@@ -361,8 +376,9 @@ class RMSNorm(nn.Module):
361
  self._use_liger = _USE_LIGER and _LigerRMSNormFn is not None
362
 
363
  def forward(self, x: torch.Tensor) -> torch.Tensor:
364
-
365
- if self._use_liger and _liger_cuda_gate(x):
 
366
  out = _liger_rmsnorm_attempt(x, self.weight, self.eps)
367
  if out is not None:
368
  return out
@@ -554,7 +570,8 @@ def _liger_fused_ce_attempt(hidden: torch.Tensor, weight: torch.Tensor, targets:
554
  try:
555
  loss_fn = _LigerFusedLinearCrossEntropyLoss(ignore_index=-1)
556
 
557
- loss = loss_fn(hidden.reshape(-1, hidden.size(-1)), weight, targets.reshape(-1))
 
558
  kernel._record_backend("liger_fused_ce_success")
559
  return loss
560
  except Exception as e: # noqa: BLE001 -- must never crash training
@@ -581,14 +598,26 @@ def chunked_cross_entropy(
581
  total_loss = torch.zeros((), device=h.device, dtype=torch.float32)
582
  n_valid = (t != ignore_index).sum().float()
583
  if n_valid == 0:
584
- return total_loss
 
 
 
 
 
 
 
 
 
585
  for start in range(0, N, chunk_size):
586
  end = min(start + chunk_size, N)
587
  h_chunk = h[start:end] # (C, D)
588
  t_chunk = t[start:end] # (C,)
589
- logits_chunk = F.linear(h_chunk, weight).float() # (C, vocab) -- freed after backward through this chunk
590
- # use sum reduction then normalise manually so partial chunks average correctly
591
- loss_chunk = F.cross_entropy(logits_chunk, t_chunk, ignore_index=ignore_index, reduction="sum")
 
 
 
592
  total_loss = total_loss + loss_chunk
593
  return total_loss / n_valid.clamp(min=1)
594
 
@@ -700,7 +729,19 @@ class DenseMHA(nn.Module):
700
  self.dropout = nn.Dropout(dropout)
701
 
702
  self.use_xsa = False # set by GPT.__init__ after reading config
703
-
 
 
 
 
 
 
 
 
 
 
 
 
704
  self._xsa_v_reps = local_num_heads // local_num_kv_heads
705
 
706
  if use_rope:
@@ -708,6 +749,17 @@ class DenseMHA(nn.Module):
708
  self.register_buffer("rope_cos", rope_cos, persistent=False)
709
  self.register_buffer("rope_sin", rope_sin, persistent=False)
710
 
 
 
 
 
 
 
 
 
 
 
 
711
  def forward(
712
  self,
713
  hidden_states: torch.Tensor,
@@ -724,6 +776,14 @@ class DenseMHA(nn.Module):
724
  k = reshape_heads(k_flat, self.num_kv_heads, self.head_dim)
725
  v = reshape_heads(v_flat, self.num_kv_heads, self.head_dim)
726
 
 
 
 
 
 
 
 
 
727
  if self.use_rope:
728
  q, k = apply_rope_qk(
729
  q,
@@ -766,15 +826,59 @@ class DenseMHA(nn.Module):
766
  )
767
  weights = None
768
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
769
  if self.use_xsa:
770
  v_xsa = v
771
  if self._xsa_v_reps > 1:
772
- v_xsa = v.repeat_interleave(self._xsa_v_reps, dim=1)
 
 
 
 
 
 
 
 
773
 
774
  v_norm = v_xsa.norm(dim=-1, keepdim=True).clamp(min=1e-6)
775
  v_unit = v_xsa / v_norm
776
  proj_coeff = (context * v_unit).sum(dim=-1, keepdim=True)
777
- context = context - proj_coeff * v_unit
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
778
 
779
  merged_context = merge_heads(context)
780
  # ─────────────────────────────────────────────────────────────────────
@@ -928,13 +1032,17 @@ class AryaSparseMoE(nn.Module):
928
  return out, self.last_aux_loss
929
 
930
  routed_out = torch.zeros_like(x_flat)
 
 
 
 
931
  for expert_id in range(self.n_routed):
932
  expert = self.routed_experts[expert_id]
933
- token_mask = (idx_flat == expert_id).any(dim=-1)
934
- if not token_mask.any():
935
- continue
936
 
937
  selected_tokens = x_flat[token_mask]
 
 
938
  expert_out = expert(selected_tokens) # Returns delta only
939
  weight = (
940
  (idx_flat[token_mask] == expert_id).float() * gate_flat[token_mask]
@@ -1128,7 +1236,10 @@ class AryaBlock(nn.Module):
1128
 
1129
  if _SIMPLE_BLOCK_PATH:
1130
  normed = self.ln_1(x)
1131
- attn_delta = module_output_tensor(self.attn(normed))
 
 
 
1132
  x = x + attn_delta
1133
  ffn_in = self.ln_2(x)
1134
  ffn_result = self.ffn(ffn_in)
@@ -1281,9 +1392,14 @@ class MTPHead(nn.Module):
1281
  self.eps = 1e-5
1282
 
1283
  def forward(self, hidden: torch.Tensor) -> torch.Tensor:
1284
- h = hidden.float()
1285
- h_norm = h * torch.rsqrt(h.pow(2).mean(-1, keepdim=True) + self.eps)
1286
- h_norm = h_norm.to(hidden.dtype) * self.weight
 
 
 
 
 
1287
  return h_norm + self.proj(h_norm)
1288
 
1289
 
@@ -1348,6 +1464,10 @@ class GPTConfig:
1348
 
1349
  use_liger: bool = _alpha_config_attr('use_liger', False)
1350
  use_xsa: bool = _alpha_config_attr('use_xsa', False) # Exclusive Self-Attention: strips self-referential component from attention output
 
 
 
 
1351
 
1352
  moe_top_k: int = _alpha_config_attr('moe_top_k', 2)
1353
 
@@ -1382,10 +1502,24 @@ class GPTConfig:
1382
 
1383
  mod_gate_entropy_coeff: float = _alpha_config_attr('mod_gate_entropy_coeff', 0.01)
1384
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1385
  class GPT(nn.Module):
1386
  """Full GPT language model"""
1387
 
1388
- def __init__(self, config):
1389
  super().__init__()
1390
  if config.vocab_size is None:
1391
  raise ValueError("config.vocab_size must be set")
@@ -1441,6 +1575,46 @@ class GPT(nn.Module):
1441
  ))
1442
  self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)
1443
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1444
 
1445
  use_mtp = bool(getattr(config, 'use_mtp', False))
1446
  mtp_depth = int(getattr(config, 'mtp_depth', 1))
@@ -1450,6 +1624,8 @@ class GPT(nn.Module):
1450
 
1451
 
1452
  self.apply(self._init_weights)
 
 
1453
 
1454
 
1455
  if config.use_xsa:
@@ -1457,6 +1633,40 @@ class GPT(nn.Module):
1457
  real_block = block.inner if isinstance(block, MoDBlock) else block
1458
  if isinstance(real_block.attn, DenseMHA):
1459
  real_block.attn.use_xsa = True
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1460
 
1461
  for block in self.transformer.h:
1462
  if isinstance(block, MoDBlock):
@@ -1587,7 +1797,12 @@ class GPT(nn.Module):
1587
  mod_gate_loss = torch.zeros((), device=x.device, dtype=torch.float32)
1588
  moe_aux_loss = torch.zeros((), device=x.device, dtype=torch.float32)
1589
 
1590
- for block, (is_mod, has_moe) in zip(self.transformer.h, self._block_flags):
 
 
 
 
 
1591
  if self.config.gradient_checkpointing and self.training:
1592
  if is_mod:
1593
  (x, moe_aux), gate_aux = torch.utils.checkpoint.checkpoint(
 
26
 
27
  _LigerRMSNormFn = None
28
  try:
29
+ # pyrefly: ignore [missing-import]
30
  from liger_kernel.ops.rms_norm import LigerRMSNormFunction as _LigerRMSNormFn
31
  except ImportError:
32
  pass
 
54
 
55
  _SIMPLE_BLOCK_PATH = not _USE_FLASH_OPS
56
 
57
+ try:
58
+ from ngram import Engram, NgramHasher
59
+ except ImportError: # pragma: no cover - fallback for package-style imports
60
+ try:
61
+ from .ngram import Engram, NgramHasher
62
+ except ImportError:
63
+ Engram = None
64
+ NgramHasher = None
65
+
66
 
67
 
68
  _alpha_config_warned: set[str] = set()
 
108
 
109
  _NVFP4_SKIP_KEYWORDS: tuple[str, ...] = (
110
  "wte", "lm_head", "ln_", "ln_f", "norm",
111
+ "router", "gate", "expert_bias", "engram",
112
  )
113
 
114
 
 
245
  return True
246
 
247
 
248
+ # torchao replaces module.weight in place and keeps the nn.Linear object, so
249
+ # identity must be tracked on the weight (same test apply_nvfp4_to_model uses).
250
+ before_ids = {fqn: id(m.weight) for fqn, m in model.named_modules() if _filter_with_gs(m, fqn)}
251
 
252
 
253
  if _INT8_AVAILABLE and _Int8Config is not None:
 
265
  _torchao_quantize(model, config_obj, filter_fn=_filter_with_gs)
266
 
267
 
268
+ current = dict(model.named_modules())
269
+ count = 0
270
+ for fqn, old_id in before_ids.items():
271
+ mod = current.get(fqn)
272
+ if mod is not None and getattr(mod, "weight", None) is not None and id(mod.weight) != old_id:
273
+ count += 1
274
  model._quant_mode = "int8"
275
  model._use_int8 = True
276
  model._use_nvfp4 = False
277
+ return count
278
 
279
  def _alpha_config_attr(name: str, default):
280
  if not hasattr(alpha_config, name):
 
376
  self._use_liger = _USE_LIGER and _LigerRMSNormFn is not None
377
 
378
  def forward(self, x: torch.Tensor) -> torch.Tensor:
379
+ # Hoist the liger CUDA guard: x.is_cuda is a tensor bool (no string
380
+ # comparison), already specialized by dynamo when _use_liger is False.
381
+ if self._use_liger and x.is_cuda:
382
  out = _liger_rmsnorm_attempt(x, self.weight, self.eps)
383
  if out is not None:
384
  return out
 
570
  try:
571
  loss_fn = _LigerFusedLinearCrossEntropyLoss(ignore_index=-1)
572
 
573
+ # LigerFusedLinearCrossEntropyLoss.forward(lin_weight, _input, target): weight FIRST.
574
+ loss = loss_fn(weight, hidden.reshape(-1, hidden.size(-1)), targets.reshape(-1))
575
  kernel._record_backend("liger_fused_ce_success")
576
  return loss
577
  except Exception as e: # noqa: BLE001 -- must never crash training
 
598
  total_loss = torch.zeros((), device=h.device, dtype=torch.float32)
599
  n_valid = (t != ignore_index).sum().float()
600
  if n_valid == 0:
601
+ # A bare zeros(()) has no grad_fn, so loss.backward() raised on an all-ignored batch.
602
+ # Multiplying by 0 keeps a graph edge to `hidden` while contributing zero gradient.
603
+ return h.float().sum() * 0.0
604
+
605
+ def _chunk_loss(h_chunk: torch.Tensor, t_chunk: torch.Tensor) -> torch.Tensor:
606
+ # Recomputed during backward via checkpoint -- only the scalar loss is
607
+ # retained, not the (chunk_size, vocab) fp32 logit tensor.
608
+ logits_chunk = F.linear(h_chunk, weight).float()
609
+ return F.cross_entropy(logits_chunk, t_chunk, ignore_index=ignore_index, reduction="sum")
610
+
611
  for start in range(0, N, chunk_size):
612
  end = min(start + chunk_size, N)
613
  h_chunk = h[start:end] # (C, D)
614
  t_chunk = t[start:end] # (C,)
615
+ # checkpoint recomputes the F.linear in backward; only the scalar
616
+ # loss_chunk is kept alive in the autograd graph between forward and
617
+ # backward, so the (C, vocab) fp32 logit tensor is never retained.
618
+ loss_chunk = torch.utils.checkpoint.checkpoint(
619
+ _chunk_loss, h_chunk, t_chunk, use_reentrant=False
620
+ )
621
  total_loss = total_loss + loss_chunk
622
  return total_loss / n_valid.clamp(min=1)
623
 
 
729
  self.dropout = nn.Dropout(dropout)
730
 
731
  self.use_xsa = False # set by GPT.__init__ after reading config
732
+ self.xsa_alpha = None # nn.Parameter, per-head gate; set by GPT.__init__ when use_xsa
733
+
734
+ self.use_qk_norm = False # set by GPT.__init__ after reading config
735
+ self.q_norm = None # RMSNorm(head_dim); set by GPT.__init__ when use_qk_norm
736
+ self.k_norm = None # RMSNorm(head_dim); set by GPT.__init__ when use_qk_norm
737
+
738
+ # Gated Attention (Qiu et al., NeurIPS 2025): per-head sigmoid gate
739
+ # on the SDPA output, before out_proj. Wired by GPT.__init__.
740
+ # Zero-init weight -> 2*sigmoid(0) = 1.0 -> gate is identity at step 0.
741
+ self.use_attn_gate = False
742
+ self.attn_gate_window = None # int: how many dims of block input to read
743
+ self.attn_gate_proj = None # nn.Linear(gate_window, num_heads, bias=False)
744
+
745
  self._xsa_v_reps = local_num_heads // local_num_kv_heads
746
 
747
  if use_rope:
 
749
  self.register_buffer("rope_cos", rope_cos, persistent=False)
750
  self.register_buffer("rope_sin", rope_sin, persistent=False)
751
 
752
+ # XSA pos-0 guard mask: registered once at init so forward() never
753
+ # allocates a fresh (T,) tensor. Non-persistent: re-created from
754
+ # max_seq_len on load, costs nothing in the state_dict.
755
+ self.register_buffer(
756
+ "_xsa_pos_mask",
757
+ torch.ones(max_seq_len),
758
+ persistent=False,
759
+ )
760
+ # Slot 0 is always masked off (position-0 has no context beyond itself).
761
+ self._xsa_pos_mask[0] = 0.0
762
+
763
  def forward(
764
  self,
765
  hidden_states: torch.Tensor,
 
776
  k = reshape_heads(k_flat, self.num_kv_heads, self.head_dim)
777
  v = reshape_heads(v_flat, self.num_kv_heads, self.head_dim)
778
 
779
+ # QK-Norm: independent of XSA, applied before the rope split. Normalizes
780
+ # q/k per head (over head_dim) to stabilize attention-logit scale, which
781
+ # keeps softmax from saturating on a few positions. Off by default --
782
+ # ablate separately from XSA so gains are attributable.
783
+ if self.use_qk_norm:
784
+ q = self.q_norm(q)
785
+ k = self.k_norm(k)
786
+
787
  if self.use_rope:
788
  q, k = apply_rope_qk(
789
  q,
 
826
  )
827
  weights = None
828
 
829
+ # ── Gated Attention (Qiu et al., NeurIPS 2025 Best Paper) ───────────────
830
+ # Per-head sigmoid gate modulates the SDPA output before XSA/out_proj.
831
+ # Gate input: first `attn_gate_window` dims of the post-norm block input
832
+ # (hidden_states). Sparse window keeps the gate head lightweight.
833
+ # 2*sigmoid(z) is used so the zero-init projection gives gate=1 (identity)
834
+ # at step 0; the head learns to suppress or amplify over training.
835
+ if self.use_attn_gate:
836
+ gw = self.attn_gate_window
837
+ gate_in = hidden_states[..., :gw] # (B, T, gw)
838
+ gate = self.attn_gate_proj(gate_in) # (B, T, num_heads)
839
+ gate = 2.0 * torch.sigmoid(gate) # (0, 2), =1.0 at init
840
+ gate = gate.transpose(1, 2).unsqueeze(-1) # (B, H, T, 1)
841
+ context = context * gate
842
+ # ────────────────────────────────────────────────────────────────────────
843
+
844
  if self.use_xsa:
845
  v_xsa = v
846
  if self._xsa_v_reps > 1:
847
+ # GQA/MQA-aware broadcast: mathematically identical to
848
+ # v.repeat_interleave(self._xsa_v_reps, dim=1) (verified via
849
+ # torch.equal), but avoids materializing a full copy of v.
850
+ Bv, Hkv, Tv, Dv = v.shape
851
+ v_xsa = (
852
+ v.reshape(Bv, Hkv, 1, Tv, Dv)
853
+ .expand(Bv, Hkv, self._xsa_v_reps, Tv, Dv)
854
+ .reshape(Bv, Hkv * self._xsa_v_reps, Tv, Dv)
855
+ )
856
 
857
  v_norm = v_xsa.norm(dim=-1, keepdim=True).clamp(min=1e-6)
858
  v_unit = v_xsa / v_norm
859
  proj_coeff = (context * v_unit).sum(dim=-1, keepdim=True)
860
+ # Learnable per-head gate (modded-nanoGPT style): tanh(alpha) in (-1, 1).
861
+ # alpha starts at 0 -> gate=0 -> forward pass is numerically identical
862
+ # to vanilla SA at step 0. Each head learns how much (if any) self-value
863
+ # component to remove, rather than a hard, deterministic subtraction.
864
+ gate = torch.tanh(self.xsa_alpha).view(1, -1, 1, 1)
865
+ correction = gate * proj_coeff * v_unit
866
+
867
+ # Position-0 guard: causal attention at position 0 has exactly one
868
+ # reachable key (itself), so context[0] == v[0] identically, at every
869
+ # step of training. XSA's projection therefore removes 100% of the
870
+ # context for that head, for any gate > 0. This is structural, not
871
+ # a training artifact -- there is no "context minus self" at
872
+ # position 0 because there is no context other than self there.
873
+ # Mask the correction off so the first token always keeps a
874
+ # full-strength value contribution.
875
+ T = context.shape[2]
876
+ # Slice the pre-built buffer; no allocation per forward. Applied for T == 1 as well:
877
+ # a 1-token input IS position 0 and must be masked exactly as it is in training.
878
+ mask = self._xsa_pos_mask[:T].to(dtype=context.dtype).view(1, 1, T, 1)
879
+ correction = correction * mask
880
+
881
+ context = context - correction
882
 
883
  merged_context = merge_heads(context)
884
  # ─────────────────────────────────────────────────────────────────────
 
1032
  return out, self.last_aux_loss
1033
 
1034
  routed_out = torch.zeros_like(x_flat)
1035
+ # Build all expert membership masks at once on the GPU with one_hot:
1036
+ # shape (B*T, n_routed), dtype bool. No .any() host syncs at all.
1037
+ # one_hot produces (B*T, top_k, n_routed); .any(dim=1) collapses top_k.
1038
+ expert_masks = F.one_hot(idx_flat, num_classes=self.n_routed).bool().any(dim=1) # (B*T, n_routed)
1039
  for expert_id in range(self.n_routed):
1040
  expert = self.routed_experts[expert_id]
1041
+ token_mask = expert_masks[:, expert_id] # (B*T,) β€” already on GPU, no sync
 
 
1042
 
1043
  selected_tokens = x_flat[token_mask]
1044
+ if selected_tokens.shape[0] == 0:
1045
+ continue
1046
  expert_out = expert(selected_tokens) # Returns delta only
1047
  weight = (
1048
  (idx_flat[token_mask] == expert_id).float() * gate_flat[token_mask]
 
1236
 
1237
  if _SIMPLE_BLOCK_PATH:
1238
  normed = self.ln_1(x)
1239
+ # In the simple path output_attentions/return_kv_cache_estimate/
1240
+ # skip_out_proj are all False, so DenseMHA returns a raw tensor --
1241
+ # no isinstance check needed, no dynamo guard on AttentionOutput.
1242
+ attn_delta = self.attn(normed)
1243
  x = x + attn_delta
1244
  ffn_in = self.ln_2(x)
1245
  ffn_result = self.ffn(ffn_in)
 
1392
  self.eps = 1e-5
1393
 
1394
  def forward(self, hidden: torch.Tensor) -> torch.Tensor:
1395
+ # F.rms_norm handles the dtype internally -- no manual fp32 upcast
1396
+ # allocation, no (B, T, D) intermediate tensor, free for any dtype.
1397
+ if _HAS_NATIVE_RMSNORM:
1398
+ h_norm = _native_rms_norm(hidden, (self.weight.shape[0],), self.weight, self.eps)
1399
+ else: # torch < 2.4: same fp32-upcast fallback as RMSNorm
1400
+ hf = hidden.float()
1401
+ h_norm = (hf / torch.sqrt(hf.pow(2).mean(-1, keepdim=True) + self.eps)).to(hidden.dtype) \
1402
+ * self.weight.to(hidden.dtype)
1403
  return h_norm + self.proj(h_norm)
1404
 
1405
 
 
1464
 
1465
  use_liger: bool = _alpha_config_attr('use_liger', False)
1466
  use_xsa: bool = _alpha_config_attr('use_xsa', False) # Exclusive Self-Attention: strips self-referential component from attention output
1467
+ use_qk_norm: bool = _alpha_config_attr('use_qk_norm', False) # RMSNorm on q/k per head, before rope; independent of use_xsa
1468
+ # Gated Attention (Qiu et al., NeurIPS 2025 Best Paper, arXiv:2505.06708)
1469
+ use_attn_gate: bool = _alpha_config_attr('use_attn_gate', False) # per-head sigmoid gate after SDPA, before out_proj
1470
+ attn_gate_window: int = _alpha_config_attr('attn_gate_window', 0) # 0 = full n_embd; 12-32 = sparse speedrun window
1471
 
1472
  moe_top_k: int = _alpha_config_attr('moe_top_k', 2)
1473
 
 
1502
 
1503
  mod_gate_entropy_coeff: float = _alpha_config_attr('mod_gate_entropy_coeff', 0.01)
1504
 
1505
+ # Engram static n-gram memory (arXiv:2601.07372). Tuples, not lists:
1506
+ # dataclass rejects mutable defaults.
1507
+ use_engram: bool = _alpha_config_attr('use_engram', False)
1508
+ engram_layer_ids: tuple = tuple(_alpha_config_attr('engram_layer_ids', (1,)))
1509
+ engram_max_ngram: int = _alpha_config_attr('engram_max_ngram', 3)
1510
+ engram_vocab_size: tuple = tuple(_alpha_config_attr('engram_vocab_size', (8192, 8192)))
1511
+ engram_embed_per_ngram: int = _alpha_config_attr('engram_embed_per_ngram', 128)
1512
+ engram_n_heads: int = _alpha_config_attr('engram_n_heads', 4)
1513
+ engram_kernel_size: int = _alpha_config_attr('engram_kernel_size', 4)
1514
+ engram_seed: int = _alpha_config_attr('engram_seed', 0)
1515
+ # Set automatically from the tokenizer lookup at build time; saved in
1516
+ # checkpoints so GPT(config) can be rebuilt without a tokenizer.
1517
+ engram_compressed_vocab: int = _alpha_config_attr('engram_compressed_vocab', 0)
1518
+
1519
  class GPT(nn.Module):
1520
  """Full GPT language model"""
1521
 
1522
+ def __init__(self, config, engram_lookup=None):
1523
  super().__init__()
1524
  if config.vocab_size is None:
1525
  raise ValueError("config.vocab_size must be set")
 
1575
  ))
1576
  self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False)
1577
 
1578
+ # ── Engram (optional) ────────────────────────────────────────────
1579
+ self.engram_hasher = None
1580
+ self.engram_layers = nn.ModuleDict()
1581
+ self._engram_layer_ids = frozenset()
1582
+ if bool(getattr(config, 'use_engram', False)):
1583
+ if Engram is None:
1584
+ raise ImportError("use_engram=True but engram.py could not be imported")
1585
+ ids = tuple(int(l) for l in config.engram_layer_ids)
1586
+ if not ids or any(not (0 <= l < config.n_layer) for l in ids) or len(set(ids)) != len(ids):
1587
+ raise ValueError(
1588
+ f"engram_layer_ids={ids} must be unique and within [0, {config.n_layer})")
1589
+ if engram_lookup is not None:
1590
+ config.engram_compressed_vocab = int(engram_lookup.max()) + 1
1591
+ if int(getattr(config, 'engram_compressed_vocab', 0)) <= 0:
1592
+ raise ValueError(
1593
+ "use_engram=True needs engram_lookup (built from the tokenizer via "
1594
+ "engram.CompressedTokenizer) or a config with engram_compressed_vocab set")
1595
+ self.engram_hasher = NgramHasher(
1596
+ layer_ids=ids,
1597
+ max_ngram=config.engram_max_ngram,
1598
+ vocab_size_per_ngram=config.engram_vocab_size,
1599
+ n_heads=config.engram_n_heads,
1600
+ compressed_vocab=config.engram_compressed_vocab,
1601
+ raw_vocab_size=config.vocab_size,
1602
+ seed=config.engram_seed,
1603
+ lookup=engram_lookup,
1604
+ )
1605
+ self.engram_layers = nn.ModuleDict({
1606
+ str(l): Engram(
1607
+ hidden_size=config.n_embd,
1608
+ head_sizes=self.engram_hasher.head_sizes[i],
1609
+ max_ngram=config.engram_max_ngram,
1610
+ embed_per_ngram=config.engram_embed_per_ngram,
1611
+ n_heads=config.engram_n_heads,
1612
+ kernel_size=config.engram_kernel_size,
1613
+ )
1614
+ for i, l in enumerate(ids)
1615
+ })
1616
+ self._engram_layer_ids = frozenset(ids)
1617
+
1618
 
1619
  use_mtp = bool(getattr(config, 'use_mtp', False))
1620
  mtp_depth = int(getattr(config, 'mtp_depth', 1))
 
1624
 
1625
 
1626
  self.apply(self._init_weights)
1627
+ for _eg in self.engram_layers.values():
1628
+ _eg.reset_parameters() # tables std=0.01 (apply() above overwrote it with 0.02)
1629
 
1630
 
1631
  if config.use_xsa:
 
1633
  real_block = block.inner if isinstance(block, MoDBlock) else block
1634
  if isinstance(real_block.attn, DenseMHA):
1635
  real_block.attn.use_xsa = True
1636
+ # Zero-init -> tanh(0) = 0 -> gate is a no-op at step 0, so the
1637
+ # forward pass starts numerically identical to vanilla SA and
1638
+ # each head learns its own correction during training.
1639
+ real_block.attn.xsa_alpha = nn.Parameter(
1640
+ torch.zeros(real_block.attn.num_heads)
1641
+ )
1642
+
1643
+ # Separate from use_xsa on purpose: QK-Norm changes attention-logit scale
1644
+ # (and therefore what XSA's projection has to work with), so it needs its
1645
+ # own on/off switch to keep the two ablatable independently.
1646
+ if config.use_qk_norm:
1647
+ for block in self.transformer.h:
1648
+ real_block = block.inner if isinstance(block, MoDBlock) else block
1649
+ if isinstance(real_block.attn, DenseMHA):
1650
+ real_block.attn.use_qk_norm = True
1651
+ real_block.attn.q_norm = RMSNorm(real_block.attn.head_dim, eps=1e-6)
1652
+ real_block.attn.k_norm = RMSNorm(real_block.attn.head_dim, eps=1e-6)
1653
+
1654
+ # Gated Attention (Qiu et al., NeurIPS 2025 Best Paper, arXiv:2505.06708).
1655
+ # Independent switch: keep ablatable from XSA and QK-Norm.
1656
+ # Gate reads the first `attn_gate_window` dims of the post-norm block input.
1657
+ # 0 (or unset) -> use full n_embd (dense gate).
1658
+ if config.use_attn_gate:
1659
+ gw = int(config.attn_gate_window) or int(config.n_embd)
1660
+ gw = min(gw, config.n_embd)
1661
+ for block in self.transformer.h:
1662
+ real_block = block.inner if isinstance(block, MoDBlock) else block
1663
+ if isinstance(real_block.attn, DenseMHA):
1664
+ real_block.attn.use_attn_gate = True
1665
+ real_block.attn.attn_gate_window = gw
1666
+ proj = nn.Linear(gw, real_block.attn.num_heads, bias=False)
1667
+ # Zero-init: 2*sigmoid(0) = 1.0 -> gate is identity at step 0.
1668
+ nn.init.zeros_(proj.weight)
1669
+ real_block.attn.attn_gate_proj = proj
1670
 
1671
  for block in self.transformer.h:
1672
  if isinstance(block, MoDBlock):
 
1797
  mod_gate_loss = torch.zeros((), device=x.device, dtype=torch.float32)
1798
  moe_aux_loss = torch.zeros((), device=x.device, dtype=torch.float32)
1799
 
1800
+ engram_hashes = self.engram_hasher(idx) if self.engram_hasher is not None else None
1801
+
1802
+ for layer_i, (block, (is_mod, has_moe)) in enumerate(zip(self.transformer.h, self._block_flags)):
1803
+ if engram_hashes is not None and layer_i in self._engram_layer_ids:
1804
+ # Official placement: h = h + Engram(h, ids) BEFORE the block's attention.
1805
+ x = x + self.engram_layers[str(layer_i)](x, engram_hashes[layer_i])
1806
  if self.config.gradient_checkpointing and self.training:
1807
  if is_mod:
1808
  (x, moe_aux), gate_aux = torch.utils.checkpoint.checkpoint(