Spaces:
Running
Running
Commit Β·
4b04e69
1
Parent(s): 507e1d8
Update model.py
Browse files
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 |
-
|
|
|
|
|
|
|
| 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 |
-
|
| 257 |
-
|
| 258 |
-
|
| 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 |
-
|
|
|
|
| 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 |
-
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 590 |
-
#
|
| 591 |
-
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 =
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 1285 |
-
|
| 1286 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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(
|