import functools import math import warnings from dataclasses import dataclass from typing import Optional, Tuple import torch import torch.nn as nn import torch.utils.checkpoint from torch.nn import functional as F _native_rms_norm = getattr(F, "rms_norm", None) _HAS_NATIVE_RMSNORM = _native_rms_norm is not None try: import config as alpha_config except ImportError: # pragma: no cover - fallback for package-style imports from . import config as alpha_config _USE_FLASH_OPS = bool(getattr(alpha_config, 'use_flash_ops', True)) _USE_FLASH_ROPE = _USE_FLASH_OPS and bool(getattr(alpha_config, 'use_flash_rope', True)) _USE_FLASH_OUTPROJ = _USE_FLASH_OPS and bool(getattr(alpha_config, 'use_flash_outproj_add_rmsnorm', True)) _USE_FLASH_SWIGLU = _USE_FLASH_OPS and bool(getattr(alpha_config, 'use_flash_swiglu', True)) _USE_FUSED_ADD_RMSNORM = bool(getattr(alpha_config, 'use_fused_add_rmsnorm', True)) _USE_LIGER = bool(getattr(alpha_config, 'use_liger', False)) _LigerRMSNormFn = None try: # pyrefly: ignore [missing-import] from liger_kernel.ops.rms_norm import LigerRMSNormFunction as _LigerRMSNormFn except ImportError: pass try: import kernel except ImportError: # pragma: no cover - fallback for package-style imports from . import kernel try: import fused_kernels except ImportError: # pragma: no cover - fallback for package-style imports try: from . import fused_kernels except ImportError: fused_kernels = None try: import flash_ops except ImportError: # pragma: no cover - fallback for package-style imports try: from . import flash_ops except ImportError: flash_ops = None _SIMPLE_BLOCK_PATH = not _USE_FLASH_OPS try: from ngram import Engram, NgramHasher except ImportError: # pragma: no cover - fallback for package-style imports try: from .ngram import Engram, NgramHasher except ImportError: Engram = None NgramHasher = None _alpha_config_warned: set[str] = set() _TORCHAO_AVAILABLE: bool = False _NVFP4_RHT_AVAILABLE: bool = False _INT8_AVAILABLE: bool = False # BUGFIX: these two were only bound inside the `try` body. If `import torchao` # succeeded but a later import in the same try block raised, they stayed # undefined and apply_int8_to_model() hit a NameError instead of the intended # RuntimeError. Bind safe defaults up front. _INT8_GROUPED_AVAILABLE: bool = False _Int8WOConfig = None # type: ignore[assignment] try: import torchao as _torchao_probe # noqa: F401 _TORCHAO_AVAILABLE = True from torchao.quantization import quantize_ as _torchao_quantize from torchao.quantization import ( Int8DynamicActivationInt8WeightConfig as _Int8Config, ) try: from torchao.quantization import Int8WeightOnlyConfig as _Int8WOConfig _INT8_GROUPED_AVAILABLE = True except ImportError: _Int8WOConfig = None # type: ignore[assignment] _INT8_GROUPED_AVAILABLE = False _INT8_AVAILABLE = True try: from torchao.prototype.mx_formats import ( NVFP4DynamicActivationNVFP4WeightConfig as _NVFP4Config, ) _NVFP4_RHT_AVAILABLE = True except ImportError: _NVFP4Config = None # type: ignore[assignment] except ImportError: _torchao_quantize = None # type: ignore[assignment] _Int8Config = None # type: ignore[assignment] _NVFP4Config = None # type: ignore[assignment] _NVFP4_DISABLED_REASON: "str | None" = None _NVFP4_SKIP_KEYWORDS: tuple[str, ...] = ( "wte", "lm_head", "ln_", "ln_f", "norm", "router", "gate", "expert_bias", "engram", ) def _nvfp4_filter(module: nn.Module, fqn: str) -> bool: """torchao filter_fn: True = quantize this module.""" if not isinstance(module, nn.Linear): return False for kw in _NVFP4_SKIP_KEYWORDS: if kw in fqn: return False out_f, in_f = module.weight.shape if in_f < 32 or out_f < 64 or in_f % 16 != 0: return False return True def _nvfp4_enabled() -> bool: global _NVFP4_DISABLED_REASON if not getattr(alpha_config, 'use_nvfp4', False): return False if not torch.cuda.is_available(): _NVFP4_DISABLED_REASON = "nvFP4: no CUDA device" return False major, _minor = torch.cuda.get_device_capability() if major < 10: msg = ( f"nvFP4 requires Blackwell (sm_100+); found sm_{major}{_minor} " f"({torch.cuda.get_device_name()}) -- disabled." ) if msg not in _alpha_config_warned: _alpha_config_warned.add(msg) warnings.warn(msg) _NVFP4_DISABLED_REASON = msg return False if not _NVFP4_RHT_AVAILABLE: _NVFP4_DISABLED_REASON = ( "nvFP4: torchao or torchao.prototype.mx_formats not available. " "Install: pip install torchao --pre" ) warnings.warn(_NVFP4_DISABLED_REASON) return False return True def apply_nvfp4_to_model(model: nn.Module) -> int: before_ids = {fqn: id(m.weight) for fqn, m in model.named_modules() if _nvfp4_filter(m, fqn)} _torchao_quantize(model, _NVFP4Config(), filter_fn=_nvfp4_filter) # BUGFIX: torchao swaps the *weight* for a tensor subclass and leaves the # module a plain nn.Linear, so the old `type(m).__name__ != "Linear"` test # always counted 0. Compare weight identity instead (same test the INT8 # path already used). current = dict(model.named_modules()) count = 0 for fqn, old_id in before_ids.items(): m = current.get(fqn) if m is not None and getattr(m, "weight", None) is not None and id(m.weight) != old_id: count += 1 model._quant_mode = "nvfp4" model._use_nvfp4 = True model._use_int8 = False return count _INT8_SKIP_KEYWORDS: tuple[str, ...] = _NVFP4_SKIP_KEYWORDS _INT8_DISABLED_REASON: "str | None" = None def _int8_filter(module: nn.Module, fqn: str) -> bool: """torchao filter_fn for INT8: True = quantize this module.""" if not isinstance(module, nn.Linear): return False for kw in _INT8_SKIP_KEYWORDS: if kw in fqn: return False out_f, in_f = module.weight.shape if in_f < 64 or out_f < 32: return False return True def _int8_enabled() -> bool: global _INT8_DISABLED_REASON if not getattr(alpha_config, 'use_int8', False): return False if not torch.cuda.is_available(): _INT8_DISABLED_REASON = "INT8: no CUDA device" return False if not _INT8_AVAILABLE: _INT8_DISABLED_REASON = ( "INT8: torchao not installed or Int8DynamicActivationInt8WeightConfig " "not found. Install: pip install torchao" ) warnings.warn(_INT8_DISABLED_REASON) return False if _nvfp4_enabled(): msg = ( "INT8 disabled: use_nvfp4=True also set and this is SM100+ hardware -- " "nvFP4 takes priority. Set use_nvfp4=False to use INT8 instead." ) if msg not in _alpha_config_warned: _alpha_config_warned.add(msg) warnings.warn(msg) _INT8_DISABLED_REASON = msg return False return True def apply_int8_to_model(model: nn.Module, group_size: int = 128) -> int: """Quantize eligible Linear layers to INT8 (Jetfire-style, dynamic activations). Config selection priority: 1. Int8DynamicActivationInt8WeightConfig() -- stable torchao, per-channel weights + dynamic per-tensor INT8 activations. Best for T4 throughput. 2. Int8WeightOnlyConfig(group_size=N) -- weight-only, activations stay fp16. Only used as fallback if (1) is unavailable (shouldn't happen on stable). Works on T4 (sm_75+). Fused kernel emitted by torch.compile via Inductor. Must be called BEFORE torch.compile(). """ def _filter_with_gs(module: nn.Module, fqn: str) -> bool: if not _int8_filter(module, fqn): return False in_f = module.weight.shape[1] if in_f % 32 != 0: # 32 = minimum alignment for any INT8 kernel return False return True # torchao replaces module.weight in place and keeps the nn.Linear object, so # identity must be tracked on the weight (same test apply_nvfp4_to_model uses). before_ids = {fqn: id(m.weight) for fqn, m in model.named_modules() if _filter_with_gs(m, fqn)} if _INT8_AVAILABLE and _Int8Config is not None: config_obj = _Int8Config() elif _INT8_GROUPED_AVAILABLE and _Int8WOConfig is not None: warnings.warn( f"[INT8] Int8DynamicActivationInt8WeightConfig unavailable -- " f"falling back to Int8WeightOnlyConfig(group_size={group_size}). " f"Activations will stay fp16 (weight-only quant).", ) config_obj = _Int8WOConfig(group_size=group_size) else: raise RuntimeError("[INT8] No usable INT8 config found in torchao -- pip install torchao") _torchao_quantize(model, config_obj, filter_fn=_filter_with_gs) current = dict(model.named_modules()) count = 0 for fqn, old_id in before_ids.items(): mod = current.get(fqn) if mod is not None and getattr(mod, "weight", None) is not None and id(mod.weight) != old_id: count += 1 model._quant_mode = "int8" model._use_int8 = True model._use_nvfp4 = False return count def _alpha_config_attr(name: str, default): if not hasattr(alpha_config, name): warnings.warn( f"config module has no attribute {name!r} -- falling back to default " f"{default!r}. If this is unexpected, check for a casing mismatch " "or rename in your config.py." ) return default return getattr(alpha_config, name) try: from liger_kernel.ops.rope import LigerRopeFunction as _LigerRopeFn except ImportError: _LigerRopeFn = None try: from liger_kernel.transformers.fused_linear_cross_entropy import ( LigerFusedLinearCrossEntropyLoss as _LigerFusedLinearCrossEntropyLoss, ) except ImportError: _LigerFusedLinearCrossEntropyLoss = None _liger_warned: set[str] = set() def _liger_warn_once(key: str, msg: str) -> None: if key not in _liger_warned: _liger_warned.add(key) warnings.warn(f"[model.py/liger] {msg}", stacklevel=3) def _liger_cuda_gate(x: torch.Tensor) -> bool: return x.device.type == "cuda" def _liger_fused_ce_cuda_gate(x: torch.Tensor) -> bool: """Liger's fused linear CE backward has a dangling-buffer bug on sm_75 (T4/Turing) under FP16 autocast + torch.compile: grad_output is read after its CUDA buffer is freed, producing cudaErrorIllegalAddress in the scalar equality check at backward line 294. Block sm_75 specifically; sm_80+ (Ampere+) is unaffected.""" if x.device.type != "cuda": return False major, minor = torch.cuda.get_device_capability(x.device) if major < 8: # sm_75 = Turing, sm_70 = Volta -- both affected return False return True def _use_liger_config() -> bool: return bool(_alpha_config_attr('use_liger', False)) def _use_flash_ops_master() -> bool: return bool(_alpha_config_attr('use_flash_ops', True)) def _use_flash_rope_config() -> bool: return _use_flash_ops_master() and bool(_alpha_config_attr('use_flash_rope', True)) def _use_flash_outproj_config() -> bool: return _use_flash_ops_master() and bool(_alpha_config_attr('use_flash_outproj_add_rmsnorm', True)) def _use_flash_swiglu_config() -> bool: return _use_flash_ops_master() and bool(_alpha_config_attr('use_flash_swiglu', True)) @torch.compiler.disable def _liger_rmsnorm_attempt(x: torch.Tensor, weight: torch.Tensor, eps: float): try: out = _LigerRMSNormFn.apply(x, weight, eps) kernel._record_backend("liger_rmsnorm_success") return out except Exception as e: # noqa: BLE001 -- must never crash training kernel._record_backend(f"liger_rmsnorm_fallback:{type(e).__name__}") _liger_warn_once( f"rmsnorm_fail_{type(e).__name__}", f"Liger RMSNorm failed on {x.device} ({e!r}) -- falling back to native/eager RMSNorm.", ) return None class RMSNorm(nn.Module): """Root Mean Square Layer Normalization (used in modern transformers)""" def __init__(self, dim, eps=1e-5): super().__init__() self.eps = eps self.weight = nn.Parameter(torch.ones(dim)) self.dim = dim self._use_liger = _USE_LIGER and _LigerRMSNormFn is not None def forward(self, x: torch.Tensor) -> torch.Tensor: # Hoist the liger CUDA guard: x.is_cuda is a tensor bool (no string # comparison), already specialized by dynamo when _use_liger is False. if self._use_liger and x.is_cuda: out = _liger_rmsnorm_attempt(x, self.weight, self.eps) if out is not None: return out if _HAS_NATIVE_RMSNORM: return _native_rms_norm(x, (self.dim,), self.weight, self.eps) x_fp32 = x.float() rms = torch.sqrt(x_fp32.pow(2).mean(-1, keepdim=True) + self.eps) weight = self.weight.to(x.dtype) return ((x_fp32 / rms).to(x.dtype)) * weight def rotate_half(x): x1, x2 = x.chunk(2, dim=-1) return torch.cat((-x2, x1), dim=-1) @dataclass class AttentionOutput: """Uniform return type so callers can optionally inspect weights / cache cost.""" output: torch.Tensor attention_weights: Optional[torch.Tensor] = None kv_cache_bytes: Optional[int] = None selected_indices: Optional[torch.Tensor] = None def module_output_tensor(x): return x.output if isinstance(x, AttentionOutput) else x def reshape_heads(x: torch.Tensor, num_heads: int, head_dim: int) -> torch.Tensor: b, t, _ = x.shape return x.view(b, t, num_heads, head_dim).transpose(1, 2) def merge_heads(x: torch.Tensor) -> torch.Tensor: b, h, t, d = x.shape return x.transpose(1, 2).contiguous().view(b, t, h * d) def expand_kv_heads(x: torch.Tensor, num_heads: int) -> torch.Tensor: b, h_kv, t, d = x.shape if h_kv == num_heads: return x if num_heads % h_kv != 0: raise ValueError(f"num_heads ({num_heads}) must be divisible by kv heads ({h_kv})") reps = num_heads // h_kv return x.repeat_interleave(reps, dim=1) @functools.lru_cache(maxsize=64) def _rope_cache(seq_len: int, dim: int, device, dtype, base: float = 10000.0): inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2, device=device, dtype=torch.float32) / dim)) t = torch.arange(seq_len, device=device, dtype=torch.float32) freqs = torch.outer(t, inv_freq) emb = torch.cat((freqs, freqs), dim=-1) return emb.cos().to(dtype), emb.sin().to(dtype) def build_rope_cache(seq_len: int, dim: int, base: float = 10000.0): inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim)) t = torch.arange(seq_len, dtype=torch.float32) freqs = torch.outer(t, inv_freq) emb = torch.cat((freqs, freqs), dim=-1) return emb.cos(), emb.sin() def apply_rope( x: torch.Tensor, positions: torch.Tensor | None = None, rope_dim: int | None = None, max_position: int | None = None, cos_full: torch.Tensor | None = None, sin_full: torch.Tensor | None = None, ) -> torch.Tensor: b, h, t, d = x.shape rope_dim = d if rope_dim is None else rope_dim is_sequential = positions is None if positions is None: positions = torch.arange(t, device=x.device) if cos_full is None or sin_full is None: if max_position is None: max_position = t max_pos = int(max_position) if max_position is not None else int(positions.max().item()) + 1 if positions.numel() > 0 else 1 cos_full, sin_full = _rope_cache(max_pos, rope_dim, x.device, x.dtype) if is_sequential: if t > cos_full.shape[0]: raise ValueError( f"RoPE table too small: highest requested position=" f"{t - 1}, table size={cos_full.shape[0]}. If this " f"happened during generate(), set gen_headroom in config.py " f"to however many extra positions past CONTEXT (block_size) " f"generation needs -- see GPT.__init__'s gen_headroom " f"wiring. This is separate from max_gen_tokens (a hard cap " f"on tokens produced per generate() call, unrelated to " f"RoPE table size). Default gen_headroom is 0 if unset." ) elif positions is not None and positions.numel() > 0: highest_pos = int(positions.max().item()) if highest_pos >= cos_full.shape[0]: raise ValueError( f"RoPE table too small: highest requested position=" f"{highest_pos}, table size={cos_full.shape[0]}. If this " f"happened during generate(), set gen_headroom in config.py " f"to however many extra positions past CONTEXT (block_size) " f"generation needs -- see GPT.__init__'s gen_headroom " f"wiring. This is separate from max_gen_tokens (a hard cap " f"on tokens produced per generate() call, unrelated to " f"RoPE table size). Default gen_headroom is 0 if unset." ) if is_sequential: cos = cos_full[:t].to(x.dtype).unsqueeze(0).unsqueeze(0) sin = sin_full[:t].to(x.dtype).unsqueeze(0).unsqueeze(0) else: cos = cos_full[positions].to(x.dtype).unsqueeze(0).unsqueeze(0) sin = sin_full[positions].to(x.dtype).unsqueeze(0).unsqueeze(0) x_rot, x_pass = x[..., :rope_dim], x[..., rope_dim:] x_rot = (x_rot * cos) + (rotate_half(x_rot) * sin) return torch.cat([x_rot, x_pass], dim=-1) if x_pass.shape[-1] > 0 else x_rot def apply_rope_qk( q: torch.Tensor, k: torch.Tensor, cos_full: torch.Tensor, sin_full: torch.Tensor, rope_dim: int, positions: torch.Tensor | None = None, ) -> tuple[torch.Tensor, torch.Tensor]: b, h, t, d = q.shape hk = k.shape[1] if (positions is None and _USE_FLASH_ROPE and flash_ops is not None and q.is_cuda): cos_t = cos_full[:t].to(q.dtype) sin_t = sin_full[:t].to(q.dtype) result = flash_ops.fused_rope_qk(q, k, cos_t, sin_t, rope_dim) if result is not None: return result if (positions is None and _USE_LIGER and _LigerRopeFn is not None and _liger_cuda_gate(q) and rope_dim == d and h == hk): result = _liger_rope_attempt(q, k, cos_full[:t], sin_full[:t]) if result is not None: return result q_rot = apply_rope(q, positions=positions, rope_dim=rope_dim, cos_full=cos_full, sin_full=sin_full) k_rot = apply_rope(k, positions=positions, rope_dim=rope_dim, cos_full=cos_full, sin_full=sin_full) return q_rot, k_rot @torch.compiler.disable def _liger_rope_attempt(q: torch.Tensor, k: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor): try: cos = cos.to(q.dtype) sin = sin.to(q.dtype) q_rot, k_rot = _LigerRopeFn.apply(q, k, cos, sin) kernel._record_backend("liger_rope_success") return q_rot, k_rot except Exception as e: # noqa: BLE001 -- must never crash training kernel._record_backend(f"liger_rope_fallback:{type(e).__name__}") _liger_warn_once( f"rope_fail_{type(e).__name__}", f"Liger RoPE failed on {q.device} ({e!r}) -- falling back to eager RoPE.", ) return None @torch.compiler.disable def _liger_fused_ce_attempt(hidden: torch.Tensor, weight: torch.Tensor, targets: torch.Tensor): try: loss_fn = _LigerFusedLinearCrossEntropyLoss(ignore_index=-1) # LigerFusedLinearCrossEntropyLoss.forward(lin_weight, _input, target): weight FIRST. loss = loss_fn(weight, hidden.reshape(-1, hidden.size(-1)), targets.reshape(-1)) kernel._record_backend("liger_fused_ce_success") return loss except Exception as e: # noqa: BLE001 -- must never crash training kernel._record_backend(f"liger_fused_ce_fallback:{type(e).__name__}") _liger_warn_once( f"fused_ce_fail_{type(e).__name__}", f"Liger fused linear CE failed on {hidden.device} ({e!r}) -- " f"falling back to eager lm_head + F.cross_entropy.", ) return None def chunked_cross_entropy( hidden: torch.Tensor, # (B, T, D) or (N, D) already flattened weight: torch.Tensor, # (vocab_size, D) targets: torch.Tensor, # (B, T) or (N,) already flattened ignore_index: int = -1, chunk_size: int = 4096, ) -> torch.Tensor: h = hidden.reshape(-1, hidden.size(-1)) # (N, D) t = targets.reshape(-1) # (N,) N = h.size(0) total_loss = torch.zeros((), device=h.device, dtype=torch.float32) n_valid = (t != ignore_index).sum().float() if n_valid == 0: # A bare zeros(()) has no grad_fn, so loss.backward() raised on an all-ignored batch. # Multiplying by 0 keeps a graph edge to `hidden` while contributing zero gradient. return h.float().sum() * 0.0 def _chunk_loss(h_chunk: torch.Tensor, t_chunk: torch.Tensor) -> torch.Tensor: # Recomputed during backward via checkpoint -- only the scalar loss is # retained, not the (chunk_size, vocab) fp32 logit tensor. logits_chunk = F.linear(h_chunk, weight).float() return F.cross_entropy(logits_chunk, t_chunk, ignore_index=ignore_index, reduction="sum") for start in range(0, N, chunk_size): end = min(start + chunk_size, N) h_chunk = h[start:end] # (C, D) t_chunk = t[start:end] # (C,) # checkpoint recomputes the F.linear in backward; only the scalar # loss_chunk is kept alive in the autograd graph between forward and # backward, so the (C, vocab) fp32 logit tensor is never retained. loss_chunk = torch.utils.checkpoint.checkpoint( _chunk_loss, h_chunk, t_chunk, use_reentrant=False ) total_loss = total_loss + loss_chunk return total_loss / n_valid.clamp(min=1) def make_causal_mask(q_len: int, k_len: int, device, q_offset: int = 0) -> torch.Tensor: q_idx = (torch.arange(q_len, device=device).view(q_len, 1) + q_offset) k_idx = torch.arange(k_len, device=device).view(1, k_len) return torch.where(k_idx <= q_idx, torch.zeros(1, device=device), torch.full((1,), float("-inf"), device=device)) def make_sliding_window_causal_mask(q_len: int, k_len: int, window_size: int, device, q_offset: int = 0) -> torch.Tensor: q_idx = (torch.arange(q_len, device=device).view(q_len, 1) + q_offset) k_idx = torch.arange(k_len, device=device).view(1, k_len) visible = (k_idx <= q_idx) & (k_idx > q_idx - window_size) return torch.where(visible, torch.zeros(1, device=device), torch.full((1,), float("-inf"), device=device)) def masked_softmax(scores: torch.Tensor, mask: torch.Tensor | None, dim: int = -1) -> torch.Tensor: if mask is not None: scores = scores + mask.to(scores.dtype) return torch.softmax(scores.float(), dim=dim).to(scores.dtype) def estimate_kv_cache_bytes( batch_size, seq_len, num_kv_heads, head_dim, dtype, compression_ratio: int = 1, window_size: int | None = None, ) -> int: elem_size = torch.zeros(1, dtype=dtype).element_size() effective_len = max(1, -(-seq_len // compression_ratio)) # ceil div if window_size is not None: effective_len = min(effective_len, window_size) return int(batch_size * num_kv_heads * effective_len * head_dim * elem_size * 2) class DenseMHA(nn.Module): def __init__( self, hidden_size: int, num_heads: int, head_dim: int | None = None, num_kv_heads: int | None = None, dropout: float = 0.0, use_rope: bool = True, causal: bool = True, bias: bool = True, max_seq_len: int = 4096, window_size: int | None = None, tp_size: int = 1, rope_dim: int | None = None, ) -> None: super().__init__() if head_dim is None: if hidden_size % num_heads != 0: raise ValueError("hidden_size must be divisible by num_heads when head_dim is omitted") head_dim = hidden_size // num_heads num_kv_heads = num_heads if num_kv_heads is None else num_kv_heads if num_heads % num_kv_heads != 0: raise ValueError(f"num_heads ({num_heads}) must be divisible by num_kv_heads ({num_kv_heads})") if tp_size > 1: raise NotImplementedError( "tensor_parallel_size > 1 is not implemented: per-rank weights are " "independently initialized (never sharded from a canonical weight), " "so DDP gradient averaging corrupts them. Set tensor_parallel_size=1. " "Real TP needs: init full weights once, slice per rank, and a " "skip_out_proj guard that honors tp_size." ) if num_heads % tp_size != 0: raise ValueError( f"tensor_parallel_size ({tp_size}) must divide num_heads ({num_heads}) " f"for this layer -- reduce tensor_parallel_size or increase n_head." ) local_num_heads = num_heads // tp_size if num_kv_heads < tp_size: local_num_kv_heads = num_kv_heads else: local_num_kv_heads = num_kv_heads // tp_size if num_kv_heads % tp_size == 0 else num_kv_heads self.hidden_size = hidden_size self.num_heads = local_num_heads self.num_kv_heads = local_num_kv_heads self.head_dim = head_dim self.inner_dim = local_num_heads * head_dim self.use_rope = use_rope self.causal = causal self.tp_size = tp_size self.tp_group = None self.window_size = window_size self.rope_dim = head_dim if rope_dim is None else rope_dim self._q_width = self.inner_dim self._k_width = self.num_kv_heads * head_dim self._v_width = self.num_kv_heads * head_dim self.qkv_proj = nn.Linear(hidden_size, self._q_width + self._k_width + self._v_width, bias=bias) self.out_proj = nn.Linear(self.inner_dim, hidden_size, bias=bias) self.dropout = nn.Dropout(dropout) self.use_xsa = False # set by GPT.__init__ after reading config self.xsa_alpha = None # nn.Parameter, per-head gate; set by GPT.__init__ when use_xsa self.use_qk_norm = False # set by GPT.__init__ after reading config self.q_norm = None # RMSNorm(head_dim); set by GPT.__init__ when use_qk_norm self.k_norm = None # RMSNorm(head_dim); set by GPT.__init__ when use_qk_norm # Gated Attention (Qiu et al., NeurIPS 2025): per-head sigmoid gate # on the SDPA output, before out_proj. Wired by GPT.__init__. # Zero-init weight -> 2*sigmoid(0) = 1.0 -> gate is identity at step 0. self.use_attn_gate = False self.attn_gate_window = None # int: how many dims of block input to read self.attn_gate_proj = None # nn.Linear(gate_window, num_heads, bias=False) self._xsa_v_reps = local_num_heads // local_num_kv_heads if use_rope: rope_cos, rope_sin = build_rope_cache(max_seq_len, self.rope_dim) self.register_buffer("rope_cos", rope_cos, persistent=False) self.register_buffer("rope_sin", rope_sin, persistent=False) # XSA pos-0 guard mask: registered once at init so forward() never # allocates a fresh (T,) tensor. Non-persistent: re-created from # max_seq_len on load, costs nothing in the state_dict. self.register_buffer( "_xsa_pos_mask", torch.ones(max_seq_len), persistent=False, ) # Slot 0 is always masked off (position-0 has no context beyond itself). self._xsa_pos_mask[0] = 0.0 def forward( self, hidden_states: torch.Tensor, output_attentions: bool = False, return_kv_cache_estimate: bool = False, skip_out_proj: bool = False, positions: torch.Tensor | None = None, ) -> torch.Tensor | AttentionOutput: batch_size, seq_len, _ = hidden_states.shape qkv = self.qkv_proj(hidden_states) q_flat, k_flat, v_flat = qkv.split([self._q_width, self._k_width, self._v_width], dim=-1) q = reshape_heads(q_flat, self.num_heads, self.head_dim) k = reshape_heads(k_flat, self.num_kv_heads, self.head_dim) v = reshape_heads(v_flat, self.num_kv_heads, self.head_dim) # QK-Norm: independent of XSA, applied before the rope split. Normalizes # q/k per head (over head_dim) to stabilize attention-logit scale, which # keeps softmax from saturating on a few positions. Off by default -- # ablate separately from XSA so gains are attributable. if self.use_qk_norm: q = self.q_norm(q) k = self.k_norm(k) if self.use_rope: q, k = apply_rope_qk( q, k, cos_full=self.rope_cos, sin_full=self.rope_sin, rope_dim=self.rope_dim, positions=positions, ) if output_attentions: if self.num_kv_heads != self.num_heads: reps = self.num_heads // self.num_kv_heads k_eager = k.repeat_interleave(reps, dim=1) v_eager = v.repeat_interleave(reps, dim=1) else: k_eager, v_eager = k, v scores = torch.matmul(q, k_eager.transpose(-2, -1)) / math.sqrt(self.head_dim) mask = None if self.window_size is not None: mask = make_sliding_window_causal_mask( seq_len, seq_len, self.window_size, hidden_states.device ).view(1, 1, seq_len, seq_len) elif self.causal: mask = make_causal_mask(seq_len, seq_len, hidden_states.device).view(1, 1, seq_len, seq_len) weights = masked_softmax(scores, mask, dim=-1) weights = self.dropout(weights) context = torch.matmul(weights, v_eager) else: context = kernel.fused_attention( q, k, v, causal=self.causal, window_size=self.window_size, dropout_p=self.dropout.p, training=self.training, ) weights = None # ── Gated Attention (Qiu et al., NeurIPS 2025 Best Paper) ─────────────── # Per-head sigmoid gate modulates the SDPA output before XSA/out_proj. # Gate input: first `attn_gate_window` dims of the post-norm block input # (hidden_states). Sparse window keeps the gate head lightweight. # 2*sigmoid(z) is used so the zero-init projection gives gate=1 (identity) # at step 0; the head learns to suppress or amplify over training. if self.use_attn_gate: gw = self.attn_gate_window gate_in = hidden_states[..., :gw] # (B, T, gw) gate = self.attn_gate_proj(gate_in) # (B, T, num_heads) gate = 2.0 * torch.sigmoid(gate) # (0, 2), =1.0 at init gate = gate.transpose(1, 2).unsqueeze(-1) # (B, H, T, 1) context = context * gate # ──────────────────────────────────────────────────────────────────────── if self.use_xsa: v_xsa = v if self._xsa_v_reps > 1: # GQA/MQA-aware broadcast: mathematically identical to # v.repeat_interleave(self._xsa_v_reps, dim=1) (verified via # torch.equal), but avoids materializing a full copy of v. Bv, Hkv, Tv, Dv = v.shape v_xsa = ( v.reshape(Bv, Hkv, 1, Tv, Dv) .expand(Bv, Hkv, self._xsa_v_reps, Tv, Dv) .reshape(Bv, Hkv * self._xsa_v_reps, Tv, Dv) ) v_norm = v_xsa.norm(dim=-1, keepdim=True).clamp(min=1e-6) v_unit = v_xsa / v_norm proj_coeff = (context * v_unit).sum(dim=-1, keepdim=True) # Learnable per-head gate (modded-nanoGPT style): tanh(alpha) in (-1, 1). # alpha starts at 0 -> gate=0 -> forward pass is numerically identical # to vanilla SA at step 0. Each head learns how much (if any) self-value # component to remove, rather than a hard, deterministic subtraction. gate = torch.tanh(self.xsa_alpha).view(1, -1, 1, 1) correction = gate * proj_coeff * v_unit # Position-0 guard: causal attention at position 0 has exactly one # reachable key (itself), so context[0] == v[0] identically, at every # step of training. XSA's projection therefore removes 100% of the # context for that head, for any gate > 0. This is structural, not # a training artifact -- there is no "context minus self" at # position 0 because there is no context other than self there. # Mask the correction off so the first token always keeps a # full-strength value contribution. T = context.shape[2] # Slice the pre-built buffer; no allocation per forward. Applied for T == 1 as well: # a 1-token input IS position 0 and must be masked exactly as it is in training. mask = self._xsa_pos_mask[:T].to(dtype=context.dtype).view(1, 1, T, 1) correction = correction * mask context = context - correction merged_context = merge_heads(context) # ───────────────────────────────────────────────────────────────────── if skip_out_proj and self.tp_size == 1 and not output_attentions and not return_kv_cache_estimate: return merged_context output = self.out_proj(merged_context) if self.tp_size > 1: if self.tp_group is None: raise RuntimeError( "DenseMHA was built with tp_size>1 but tp_group was never " "set -- call GPT.set_tp_group(group) after construction, " "before the first forward pass." ) torch.distributed.all_reduce(output, op=torch.distributed.ReduceOp.SUM, group=self.tp_group) if output_attentions or return_kv_cache_estimate: cache_bytes = None if return_kv_cache_estimate: cache_bytes = estimate_kv_cache_bytes( batch_size, seq_len, self.num_kv_heads, self.head_dim, hidden_states.dtype, window_size=self.window_size, ) return AttentionOutput(output=output, attention_weights=weights if output_attentions else None, kv_cache_bytes=cache_bytes) return output class SwiGLU_FFN(nn.Module): """ SwiGLU Feed Forward Network. """ def __init__(self, d_model, ffn_mult=8/3, dropout=0.0, tp_size: int = 1): super().__init__() raw_hidden = d_model * ffn_mult hidden_dim = int(((raw_hidden + 63) // 64) * 64) if hidden_dim % tp_size != 0: raise ValueError( f"tensor_parallel_size ({tp_size}) must divide the FFN hidden_dim " f"({hidden_dim} = round({d_model} * {ffn_mult}) -> next mult of 64) " f"-- adjust ffn_mult or tensor_parallel_size." ) local_hidden_dim = hidden_dim // tp_size self.tp_size = tp_size self.tp_group = None # set post-construction via GPT.set_tp_group() self.W_gate = nn.Linear(d_model, local_hidden_dim, bias=False) self.W_value = nn.Linear(d_model, local_hidden_dim, bias=False) self.W_out = nn.Linear(local_hidden_dim, d_model, bias=False) self.dropout = nn.Dropout(dropout) self._use_flash_swiglu = _USE_FLASH_SWIGLU and flash_ops is not None def forward(self, x): gate_raw = self.W_gate(x) value = self.W_value(x) if _SIMPLE_BLOCK_PATH or not self._use_flash_swiglu: hidden = F.silu(gate_raw) * value else: hidden = flash_ops.fused_swiglu(gate_raw, value) if gate_raw.is_cuda else F.silu(gate_raw) * value hidden = self.dropout(hidden) out = self.W_out(hidden) if self.tp_size > 1: if self.tp_group is None: raise RuntimeError( "SwiGLU_FFN was built with tp_size>1 but tp_group was never " "set -- call GPT.set_tp_group(group) after construction, " "before the first forward pass." ) torch.distributed.all_reduce(out, op=torch.distributed.ReduceOp.SUM, group=self.tp_group) return out # Return delta only class AryaSparseMoE(nn.Module): """ DeepSeek-style Sparse Mixture of Experts. Expects input to be ALREADY normalized (Pre-LN). Returns the combined FFN delta. """ def __init__(self, d_model, n_experts=8, n_shared=1, top_k=2, ffn_mult=8/3, dropout=0.0, use_capacity_routing=False, capacity_factor=1.25, fixed_capacity: int | None = None): super().__init__() self.n_experts = n_experts self.n_routed = n_experts - n_shared self.top_k = top_k self.d_model = d_model self.use_capacity_routing = use_capacity_routing self.capacity_factor = capacity_factor self.fixed_capacity = fixed_capacity if self.n_routed < top_k: raise ValueError( f"AryaSparseMoE config error: n_experts={n_experts} - n_shared={n_shared} " f"= {self.n_routed} routed experts, but top_k={top_k} routing needs at " f"least {top_k} routed experts to choose from." ) self.shared_expert = SwiGLU_FFN(d_model, ffn_mult, dropout=dropout) self.routed_experts = nn.ModuleList([ SwiGLU_FFN(d_model, ffn_mult, dropout=dropout) for _ in range(self.n_routed) ]) self.router = nn.Linear(d_model, self.n_routed, bias=False) self.expert_bias = nn.Parameter(torch.zeros(self.n_routed), requires_grad=False) self.aux_loss_weight = float(getattr(alpha_config, 'moe_aux_loss_weight', 0.01)) self.last_aux_loss = torch.zeros((), dtype=torch.float32) def forward(self, x): B, T, D = x.shape shared_out = self.shared_expert(x) router_logits = self.router(x) if self.training and torch.is_grad_enabled(): biased_logits = router_logits.float() + self.expert_bias.float() _, topk_indices = biased_logits.topk(self.top_k, dim=-1) full_probs = F.softmax(router_logits.float(), dim=-1) topk_probs = full_probs.gather(-1, topk_indices).to(router_logits.dtype) topk_probs = topk_probs / topk_probs.sum(dim=-1, keepdim=True).clamp_min(1e-9) one_hot = F.one_hot(topk_indices, num_classes=self.n_routed).float() f = one_hot.mean(dim=(0, 1)) P = full_probs.mean(dim=(0, 1)) self.last_aux_loss = self.aux_loss_weight * self.n_routed * (f * P).sum() # Do not mutate expert_bias here. Checkpoint recomputation runs # this forward a second time, so an in-forward update would apply # twice per optimizer step and make routing state nondeterministic. else: full_probs = F.softmax(router_logits.float(), dim=-1) topk_probs, topk_indices = full_probs.topk(self.top_k, dim=-1) topk_probs = (topk_probs / topk_probs.sum(dim=-1, keepdim=True).clamp_min(1e-9)).to(router_logits.dtype) self.last_aux_loss = torch.zeros((), device=x.device, dtype=torch.float32) gate_weights = topk_probs x_flat = x.view(B * T, D) gate_flat = gate_weights.view(B * T, self.top_k) idx_flat = topk_indices.view(B * T, self.top_k) if self.use_capacity_routing: routed_out = self._forward_capacity_routed(x_flat, gate_flat, idx_flat) routed_out = routed_out.view(B, T, D) out = shared_out + routed_out return out, self.last_aux_loss routed_out = torch.zeros_like(x_flat) # Build all expert membership masks at once on the GPU with one_hot: # shape (B*T, n_routed), dtype bool. No .any() host syncs at all. # one_hot produces (B*T, top_k, n_routed); .any(dim=1) collapses top_k. expert_masks = F.one_hot(idx_flat, num_classes=self.n_routed).bool().any(dim=1) # (B*T, n_routed) for expert_id in range(self.n_routed): expert = self.routed_experts[expert_id] token_mask = expert_masks[:, expert_id] # (B*T,) — already on GPU, no sync selected_tokens = x_flat[token_mask] if selected_tokens.shape[0] == 0: continue expert_out = expert(selected_tokens) # Returns delta only weight = ( (idx_flat[token_mask] == expert_id).float() * gate_flat[token_mask] ).sum(dim=-1) routed_out[token_mask] += weight.unsqueeze(-1) * expert_out routed_out = routed_out.view(B, T, D) out = shared_out + routed_out return out, self.last_aux_loss def _forward_capacity_routed(self, x_flat, gate_flat, idx_flat): n_tokens, D = x_flat.shape n_routed, capacity = self.n_routed, self._capacity(n_tokens) flat_expert_ids = idx_flat.reshape(-1) flat_gates = gate_flat.reshape(-1) flat_token_ids = torch.arange(n_tokens, device=x_flat.device).unsqueeze(1).expand(-1, self.top_k).reshape(-1) one_hot = F.one_hot(flat_expert_ids, num_classes=n_routed).to(torch.float32) position_in_expert = (one_hot.cumsum(dim=0) * one_hot).sum(dim=1).long() - 1 keep = position_in_expert < capacity safe_position = position_in_expert.clamp(min=0, max=capacity - 1) flat_slot = flat_expert_ids * capacity + safe_position gathered_x = x_flat[flat_token_ids] * keep.unsqueeze(-1).to(x_flat.dtype) dispatch = torch.zeros(n_routed * capacity, D, device=x_flat.device, dtype=x_flat.dtype) dispatch.index_add_(0, flat_slot, gathered_x) dispatch = dispatch.view(n_routed, capacity, D) expert_out = torch.stack( [expert(dispatch[e]) for e, expert in enumerate(self.routed_experts)], dim=0, ).view(n_routed * capacity, D) contrib = expert_out[flat_slot] * keep.unsqueeze(-1).to(x_flat.dtype) * flat_gates.unsqueeze(-1) routed_out = torch.zeros(n_tokens, D, device=x_flat.device, dtype=x_flat.dtype) routed_out.index_add_(0, flat_token_ids, contrib) return routed_out def _capacity(self, n_tokens) -> int: if self.fixed_capacity is not None: return self.fixed_capacity total_assignments = n_tokens * self.top_k avg_load = -(-total_assignments // self.n_routed) numer = int(round(self.capacity_factor * 1000)) capacity = -(-(avg_load * numer) // 1000) return capacity + 1 def build_attention_layers( num_layers: int, hidden_size: int, num_heads: int, dropout: float = 0.0, bias: bool = True, num_kv_heads: int | None = None, pattern: str = "dense", max_seq_len: int = 4096, sliding_window_size: int = 256, tp_size: int = 1, rope_dim: int | None = None, ): if num_kv_heads is None: num_kv_heads = num_heads layers = [] layer_types = [] if pattern == "dense": for _ in range(num_layers): layers.append(DenseMHA( hidden_size=hidden_size, num_heads=num_heads, num_kv_heads=num_kv_heads, dropout=dropout, bias=bias, max_seq_len=max_seq_len, window_size=None, tp_size=tp_size, rope_dim=rope_dim, )) layer_types.append("dense") elif pattern == "pyramid_swa": if num_kv_heads <= 1: pyramid_kv_heads = [1, 1, 1] else: import math as _math mqa_floor = max(2, _math.ceil(num_kv_heads / 4)) mqa_floor = min(mqa_floor, num_kv_heads) # never exceed full heads pyramid_kv_heads = sorted(set([ mqa_floor, max(mqa_floor, _math.ceil(num_kv_heads / 2)), num_kv_heads, ])) while len(pyramid_kv_heads) < 3: pyramid_kv_heads.append(num_kv_heads) pyramid_kv_heads = pyramid_kv_heads[:3] for i in range(num_layers): pos_in_group = i % 4 is_last_layer = (i == num_layers - 1) if pos_in_group == 3 or is_last_layer: layers.append(DenseMHA( hidden_size=hidden_size, num_heads=num_heads, num_kv_heads=num_kv_heads, dropout=dropout, bias=bias, max_seq_len=max_seq_len, window_size=None, tp_size=tp_size, rope_dim=rope_dim, )) layer_types.append("global") else: kv_heads_here = pyramid_kv_heads[pos_in_group] layers.append(DenseMHA( hidden_size=hidden_size, num_heads=num_heads, num_kv_heads=kv_heads_here, dropout=dropout, bias=bias, max_seq_len=max_seq_len, window_size=sliding_window_size, tp_size=tp_size, rope_dim=rope_dim, )) layer_types.append("swa") elif pattern == "mod": for i in range(num_layers): layers.append(DenseMHA( hidden_size=hidden_size, num_heads=num_heads, num_kv_heads=num_kv_heads, dropout=dropout, bias=bias, max_seq_len=max_seq_len, window_size=None, tp_size=tp_size, rope_dim=rope_dim, )) is_full_dense = (i % 4 == 3) or (i == num_layers - 1) layer_types.append("mod_dense" if is_full_dense else "mod") else: raise ValueError( f"Unknown attention pattern {pattern!r} -- supported patterns: " f"'dense', 'pyramid_swa', 'mod'. " f"(BigBird/conv-hybrid/LSA/HCA patterns have been removed.)" ) return nn.ModuleList(layers), layer_types class AryaBlock(nn.Module): """Single transformer block: Layer Norm → Attention → MLP (with residual connections). `attn_module` is DenseMHA, pre-built by GPT via build_attention_layers. Every layer is fully standalone -- no cross-layer attention/KV sharing of any kind. `role` is purely a label ("dense" | "lsa") for inspection/debugging. """ def __init__(self, config, attn_module: nn.Module, role: str = "dense"): super().__init__() self.ln_1 = RMSNorm(config.n_embd, eps=1e-5) # Pre-LN for the MLP branch: ensure router and experts see the # identical normalized tensor the MLPs expect. self.ln_2 = RMSNorm(config.n_embd, eps=1e-5) self.attn = attn_module self.role = role # "dense" | "lsa" self._use_fused_outproj = ( _USE_FLASH_OUTPROJ and flash_ops is not None and self.attn.out_proj.bias is None ) self.use_moe = getattr(config, 'use_moe', False) if self.use_moe: if getattr(config, 'tp_size', 1) > 1: raise ValueError( "tensor_parallel_size > 1 is not supported together with use_moe=True -- " "MoE needs its own (unimplemented here) expert-parallel sharding scheme, " "column/row-parallel head splitting does not apply to expert routing. " "Set use_moe=False or tensor_parallel_size=1." ) self.ffn = AryaSparseMoE( d_model=config.n_embd, n_experts=config.n_experts, n_shared=config.n_shared, top_k=getattr(config, 'moe_top_k', 2), ffn_mult=config.ffn_mult, dropout=config.dropout, use_capacity_routing=False, fixed_capacity=getattr(config, 'moe_fixed_capacity', None), ) else: self.ffn = SwiGLU_FFN(config.n_embd, config.ffn_mult, dropout=config.dropout, tp_size=getattr(config, 'tp_size', 1)) def forward(self, x): """Every attention layer (DenseMHA, LSA) is standalone and computes its own K/V from `x` alone -- no external_kv, no cross-layer sharing, no produced_kv to pass on. """ if _SIMPLE_BLOCK_PATH: normed = self.ln_1(x) # In the simple path output_attentions/return_kv_cache_estimate/ # skip_out_proj are all False, so DenseMHA returns a raw tensor -- # no isinstance check needed, no dynamo guard on AttentionOutput. attn_delta = self.attn(normed) x = x + attn_delta ffn_in = self.ln_2(x) ffn_result = self.ffn(ffn_in) if self.use_moe: ffn_delta, moe_aux = ffn_result return x + ffn_delta, moe_aux else: return x + ffn_result normed = self.ln_1(x) if self._use_fused_outproj: fused_out = None raw_ctx = None attn_delta = None attn_raw = self.attn(normed, skip_out_proj=True) candidate = module_output_tensor(attn_raw) if candidate.is_cuda and candidate.shape[-1] == self.attn.out_proj.weight.shape[1]: raw_ctx = candidate gemm_dtype = raw_ctx.dtype fused_out = flash_ops.fused_outproj_add_rmsnorm( raw_ctx.reshape(-1, raw_ctx.shape[-1]), self.attn.out_proj.weight.to(dtype=gemm_dtype), x.reshape(-1, x.shape[-1]).to(dtype=gemm_dtype), self.ln_2.weight.to(dtype=gemm_dtype), self.ln_2.eps, ) if fused_out is not None: residual_out, normed_out = fused_out x = residual_out.to(dtype=x.dtype).reshape(x.shape) ffn_in = normed_out.reshape(x.shape) else: skip_out_proj_honored = self.attn.tp_size == 1 if skip_out_proj_honored: raw_ctx = candidate attn_delta = self.attn.out_proj(raw_ctx) else: attn_delta = candidate if fused_out is None: if attn_delta is None: # Fused kernel returned None despite matching shape # (rare fallback) -- raw_ctx was already computed above. attn_delta = self.attn.out_proj(raw_ctx) tier2_out = None if _USE_FUSED_ADD_RMSNORM and fused_kernels is not None: tier2_out = fused_kernels.fused_add_rmsnorm( x, attn_delta, self.ln_2.weight.to(dtype=x.dtype), self.ln_2.eps ) if tier2_out is not None: x, ffn_in = tier2_out else: # Tier 3: fully eager, byte-for-byte what this block # always ran before any fusion work this session. x = x + attn_delta ffn_in = self.ln_2(x) else: # No fused out-proj path available/enabled: single eager path, # no dead branch checks (fused_out/raw_ctx/attn_delta are never # anything but this one shape when _use_fused_outproj is False). attn_raw = self.attn(normed) attn_delta = module_output_tensor(attn_raw) tier2_out = None if _USE_FUSED_ADD_RMSNORM and fused_kernels is not None: tier2_out = fused_kernels.fused_add_rmsnorm( x, attn_delta, self.ln_2.weight.to(dtype=x.dtype), self.ln_2.eps ) if tier2_out is not None: x, ffn_in = tier2_out else: # Tier 3: fully eager, byte-for-byte what this block # always ran before any fusion work this session. x = x + attn_delta ffn_in = self.ln_2(x) ffn_result = self.ffn(ffn_in) if self.use_moe: ffn_delta, moe_aux = ffn_result x = x + ffn_delta return x, moe_aux else: return x + ffn_result class MoDBlock(nn.Module): def __init__(self, config, inner: nn.Module, capacity: float = 0.125): super().__init__() self.inner = inner self.capacity = capacity self.keep_target = float(capacity) self._budget_coeff = float(getattr(config, 'mod_budget_coeff', 0.05)) self._entropy_coeff = float(getattr(config, 'mod_gate_entropy_coeff', 0.0)) self.hard_eval = bool(getattr(config, 'mod_hard_eval', False)) self._router_norm = RMSNorm(config.n_embd, eps=1e-5) self.router = nn.Linear(config.n_embd, 1, bias=True) self.last_gate_aux = None def forward(self, x): gate = torch.sigmoid(self.router(self._router_norm(x))) aux = torch.zeros((), device=x.device, dtype=torch.float32) if self.training: g = gate.float() if self._budget_coeff > 0.0: aux = aux + self._budget_coeff * (g.mean() - self.keep_target) ** 2 if self._entropy_coeff > 0.0: gc = g.clamp(1e-6, 1.0 - 1e-6) aux = aux + self._entropy_coeff * ( -(gc * gc.log() + (1 - gc) * (1 - gc).log()).mean() ) # BUGFIX: `aux` must stay attached to the graph -- returning # aux.detach() made the MoD budget/entropy regularizer contribute # exactly zero gradient to the router, i.e. it was inert. self.last_gate_aux = aux.detach() # detached copy, for logging only else: aux = torch.zeros((), device=x.device, dtype=torch.float32) self.last_gate_aux = aux if self.inner.use_moe: inner_out, moe_aux = self.inner(x) else: inner_out = self.inner(x) moe_aux = torch.zeros((), device=x.device, dtype=torch.float32) out = x + gate * (inner_out - x) return (out, moe_aux), aux def extra_repr(self) -> str: return ( f"gate_based=True,keep_target={self.keep_target}," f"budget_coeff={self._budget_coeff},entropy_coeff={self._entropy_coeff}" ) class MTPHead(nn.Module): """Residual shared-head auxiliary for a future-token prediction depth.""" def __init__(self, hidden_size: int, ln_f_weight: torch.Tensor): super().__init__() self.weight = nn.Parameter(ln_f_weight.detach().clone()) self.proj = nn.Linear(hidden_size, hidden_size, bias=False) nn.init.zeros_(self.proj.weight) self.eps = 1e-5 def forward(self, hidden: torch.Tensor) -> torch.Tensor: # F.rms_norm handles the dtype internally -- no manual fp32 upcast # allocation, no (B, T, D) intermediate tensor, free for any dtype. if _HAS_NATIVE_RMSNORM: h_norm = _native_rms_norm(hidden, (self.weight.shape[0],), self.weight, self.eps) else: # torch < 2.4: same fp32-upcast fallback as RMSNorm hf = hidden.float() h_norm = (hf / torch.sqrt(hf.pow(2).mean(-1, keepdim=True) + self.eps)).to(hidden.dtype) \ * self.weight.to(hidden.dtype) return h_norm + self.proj(h_norm) class MultiTokenPrediction(nn.Module): def __init__(self, config, depth: int = 1, ln_f_weight: torch.Tensor | None = None): super().__init__() self.depth = depth initial_weight = ( ln_f_weight if ln_f_weight is not None else torch.ones(config.n_embd) ) self.heads = nn.ModuleList([ MTPHead(config.n_embd, initial_weight) for _ in range(depth) ]) def compute_loss( self, x: torch.Tensor, # (B, T, D) -- pre-lm_head hidden states targets: torch.Tensor, # (B, T) -- already the +1 shifted targets lm_head_weight: torch.Tensor, # (vocab_size, D) -- tied weight mtp_lambda: float = 0.3, ) -> torch.Tensor: total = torch.zeros((), device=x.device, dtype=torch.float32) valid_depths = 0 for k, head in enumerate(self.heads): offset = k + 1 # additional token offset beyond the main +1 shift Tk = x.size(1) - offset if Tk <= 0: continue h = head(x[:, :Tk]) # (B, Tk, D) tgt_k = targets[:, offset:] # (B, Tk) # Keep sum reduction: all-ignored chunks must contribute zero, not NaN. loss_k = chunked_cross_entropy(h, lm_head_weight, tgt_k, ignore_index=-1) total = total + loss_k valid_depths += 1 if valid_depths == 0: return torch.zeros((), device=x.device, dtype=torch.float32) return mtp_lambda * (total / valid_depths) @dataclass class GPTConfig: """Configuration for GPT model.""" block_size: int = _alpha_config_attr('CONTEXT', 1024) vocab_size: int = _alpha_config_attr('vocab_size', 50304) n_layer: int = _alpha_config_attr('numberoflayers', 12) n_head: int = _alpha_config_attr('numberofheads', 12) n_embd: int = _alpha_config_attr('D_MODEL', 768) dropout: float = _alpha_config_attr('dropout', 0.0) bias: bool = _alpha_config_attr('bias', False) ffn_mult: float = _alpha_config_attr('ffn_mult', 4.0) d_rope: int | None = _alpha_config_attr('d_rope', None) top_k: int = _alpha_config_attr('top_k', 64) n_experts: int = _alpha_config_attr('n_experts', 8) n_shared: int = _alpha_config_attr('n_shared', 1) use_moe: bool = _alpha_config_attr('use_moe', False) use_liger: bool = _alpha_config_attr('use_liger', False) use_xsa: bool = _alpha_config_attr('use_xsa', False) # Exclusive Self-Attention: strips self-referential component from attention output use_qk_norm: bool = _alpha_config_attr('use_qk_norm', False) # RMSNorm on q/k per head, before rope; independent of use_xsa # Gated Attention (Qiu et al., NeurIPS 2025 Best Paper, arXiv:2505.06708) use_attn_gate: bool = _alpha_config_attr('use_attn_gate', False) # per-head sigmoid gate after SDPA, before out_proj attn_gate_window: int = _alpha_config_attr('attn_gate_window', 0) # 0 = full n_embd; 12-32 = sparse speedrun window moe_top_k: int = _alpha_config_attr('moe_top_k', 2) gradient_checkpointing: bool = _alpha_config_attr('gradient_checkpointing', False) # trade ~20-30% more compute for substantially less activation memory (grows with n_layer -- see GPT.forward) num_kv_heads: int | None = _alpha_config_attr('num_kv_heads', None) pattern: str = _alpha_config_attr('pattern', 'pyramid_swa') sliding_window_size: int = _alpha_config_attr('sliding_window_size', 256) tp_size: int = _alpha_config_attr('tensor_parallel_size', 1) gen_headroom: int = _alpha_config_attr('gen_headroom', 0) max_gen_tokens: int | None = _alpha_config_attr('max_gen_tokens', None) mod_capacity: float = _alpha_config_attr('mod_capacity', 0.125) use_mtp: bool = _alpha_config_attr('use_mtp', False) mtp_depth: int = _alpha_config_attr('mtp_depth', 1) mtp_lambda: float = _alpha_config_attr('mtp_lambda', 0.3) use_nvfp4: bool = _alpha_config_attr('use_nvfp4', False) use_int8: bool = _alpha_config_attr('use_int8', False) int8_group_size: int = _alpha_config_attr('int8_group_size', 128) mod_gate_entropy_coeff: float = _alpha_config_attr('mod_gate_entropy_coeff', 0.01) # Engram static n-gram memory (arXiv:2601.07372). Tuples, not lists: # dataclass rejects mutable defaults. use_engram: bool = _alpha_config_attr('use_engram', False) engram_layer_ids: tuple = tuple(_alpha_config_attr('engram_layer_ids', (1,))) engram_max_ngram: int = _alpha_config_attr('engram_max_ngram', 3) engram_vocab_size: tuple = tuple(_alpha_config_attr('engram_vocab_size', (8192, 8192))) engram_embed_per_ngram: int = _alpha_config_attr('engram_embed_per_ngram', 128) engram_n_heads: int = _alpha_config_attr('engram_n_heads', 4) engram_kernel_size: int = _alpha_config_attr('engram_kernel_size', 4) engram_seed: int = _alpha_config_attr('engram_seed', 0) # Set automatically from the tokenizer lookup at build time; saved in # checkpoints so GPT(config) can be rebuilt without a tokenizer. engram_compressed_vocab: int = _alpha_config_attr('engram_compressed_vocab', 0) class GPT(nn.Module): """Full GPT language model""" def __init__(self, config, engram_lookup=None): super().__init__() if config.vocab_size is None: raise ValueError("config.vocab_size must be set") if config.block_size is None: raise ValueError("config.block_size must be set") self.config = config head_dim = config.n_embd // config.n_head rope_dim = head_dim if getattr(config, 'd_rope', None) is None else config.d_rope gen_headroom = max(0, int(getattr(config, 'gen_headroom', 0))) rope_table_len = config.block_size + gen_headroom attn_layers, self.layer_types = build_attention_layers( num_layers=config.n_layer, hidden_size=config.n_embd, num_heads=config.n_head, dropout=config.dropout, bias=config.bias, num_kv_heads=config.num_kv_heads, pattern=config.pattern, max_seq_len=rope_table_len, sliding_window_size=config.sliding_window_size, tp_size=config.tp_size, rope_dim=rope_dim, ) mod_capacity = float(getattr(config, 'mod_capacity', 0.125)) blocks = [] for i in range(config.n_layer): block = AryaBlock(config, attn_layers[i], role=self.layer_types[i]) if self.layer_types[i] == "mod": block = MoDBlock(config, block, capacity=mod_capacity) blocks.append(block) # Hoist once: avoid an isinstance(block, MoDBlock) check (and the # use_moe lookup that decides whether a block yields a moe_aux # tensor) on every block on every forward. self._block_flags = [ (isinstance(b, MoDBlock), b.inner.use_moe if isinstance(b, MoDBlock) else b.use_moe) for b in blocks ] self.transformer = nn.ModuleDict(dict( wte=nn.Embedding(config.vocab_size, config.n_embd), drop=nn.Dropout(config.dropout), h=nn.ModuleList(blocks), ln_f=RMSNorm(config.n_embd, eps=1e-5), )) self.lm_head = nn.Linear(config.n_embd, config.vocab_size, bias=False) # ── Engram (optional) ──────────────────────────────────────────── self.engram_hasher = None self.engram_layers = nn.ModuleDict() self._engram_layer_ids = frozenset() if bool(getattr(config, 'use_engram', False)): if Engram is None: raise ImportError("use_engram=True but engram.py could not be imported") ids = tuple(int(l) for l in config.engram_layer_ids) if not ids or any(not (0 <= l < config.n_layer) for l in ids) or len(set(ids)) != len(ids): raise ValueError( f"engram_layer_ids={ids} must be unique and within [0, {config.n_layer})") if engram_lookup is not None: config.engram_compressed_vocab = int(engram_lookup.max()) + 1 if int(getattr(config, 'engram_compressed_vocab', 0)) <= 0: raise ValueError( "use_engram=True needs engram_lookup (built from the tokenizer via " "engram.CompressedTokenizer) or a config with engram_compressed_vocab set") self.engram_hasher = NgramHasher( layer_ids=ids, max_ngram=config.engram_max_ngram, vocab_size_per_ngram=config.engram_vocab_size, n_heads=config.engram_n_heads, compressed_vocab=config.engram_compressed_vocab, raw_vocab_size=config.vocab_size, seed=config.engram_seed, lookup=engram_lookup, ) self.engram_layers = nn.ModuleDict({ str(l): Engram( hidden_size=config.n_embd, head_sizes=self.engram_hasher.head_sizes[i], max_ngram=config.engram_max_ngram, embed_per_ngram=config.engram_embed_per_ngram, n_heads=config.engram_n_heads, kernel_size=config.engram_kernel_size, ) for i, l in enumerate(ids) }) self._engram_layer_ids = frozenset(ids) use_mtp = bool(getattr(config, 'use_mtp', False)) mtp_depth = int(getattr(config, 'mtp_depth', 1)) self.mtp = MultiTokenPrediction(config, depth=mtp_depth) if use_mtp else None self._mtp_lambda = float(getattr(config, 'mtp_lambda', 0.3)) self.apply(self._init_weights) for _eg in self.engram_layers.values(): _eg.reset_parameters() # tables std=0.01 (apply() above overwrote it with 0.02) if config.use_xsa: for block in self.transformer.h: real_block = block.inner if isinstance(block, MoDBlock) else block if isinstance(real_block.attn, DenseMHA): real_block.attn.use_xsa = True # Zero-init -> tanh(0) = 0 -> gate is a no-op at step 0, so the # forward pass starts numerically identical to vanilla SA and # each head learns its own correction during training. real_block.attn.xsa_alpha = nn.Parameter( torch.zeros(real_block.attn.num_heads) ) # Separate from use_xsa on purpose: QK-Norm changes attention-logit scale # (and therefore what XSA's projection has to work with), so it needs its # own on/off switch to keep the two ablatable independently. if config.use_qk_norm: for block in self.transformer.h: real_block = block.inner if isinstance(block, MoDBlock) else block if isinstance(real_block.attn, DenseMHA): real_block.attn.use_qk_norm = True real_block.attn.q_norm = RMSNorm(real_block.attn.head_dim, eps=1e-6) real_block.attn.k_norm = RMSNorm(real_block.attn.head_dim, eps=1e-6) # Gated Attention (Qiu et al., NeurIPS 2025 Best Paper, arXiv:2505.06708). # Independent switch: keep ablatable from XSA and QK-Norm. # Gate reads the first `attn_gate_window` dims of the post-norm block input. # 0 (or unset) -> use full n_embd (dense gate). if config.use_attn_gate: gw = int(config.attn_gate_window) or int(config.n_embd) gw = min(gw, config.n_embd) for block in self.transformer.h: real_block = block.inner if isinstance(block, MoDBlock) else block if isinstance(real_block.attn, DenseMHA): real_block.attn.use_attn_gate = True real_block.attn.attn_gate_window = gw proj = nn.Linear(gw, real_block.attn.num_heads, bias=False) # Zero-init: 2*sigmoid(0) = 1.0 -> gate is identity at step 0. nn.init.zeros_(proj.weight) real_block.attn.attn_gate_proj = proj for block in self.transformer.h: if isinstance(block, MoDBlock): with torch.no_grad(): block.router.bias.fill_(2.0) if self.mtp is not None: with torch.no_grad(): for head in self.mtp.heads: head.weight.copy_(self.transformer.ln_f.weight) head.proj.weight.zero_() self.transformer.wte.weight = self.lm_head.weight residual_std = 0.02 / math.sqrt(config.n_layer) for name, p in self.named_parameters(): if name.endswith("out_proj.weight") or name.endswith("W_out.weight"): torch.nn.init.normal_(p, mean=0.0, std=residual_std) self._use_nvfp4: bool = False self._use_int8: bool = False self._quant_mode = None if _nvfp4_enabled(): n_quantized = apply_nvfp4_to_model(self) self._use_nvfp4 = True self._quant_mode = "nvfp4" warnings.warn( f"[nvFP4/RHT] torchao quantized {n_quantized} Linear layers " f"(Blackwell SM100+). Weight: FP4 e2m1 per-block-16 RHT. " f"Activation: FP4 dynamic per-tensor scale (FP8 intermediate). " f"Fused GEMM kernel emitted by torch.compile. " f"Skipped: wte, lm_head, norms, router, gate, small/misaligned linears. " f"Verify val_loss vs bf16 baseline before committing to a long run.", stacklevel=2, ) elif _int8_enabled(): gs = int(getattr(config, 'int8_group_size', 128)) n_quantized = apply_int8_to_model(self, group_size=gs) self._use_int8 = True self._quant_mode = "int8" warnings.warn( f"[INT8/Jetfire] torchao quantized {n_quantized} Linear layers " f"(T4/Turing+). Weight: INT8 per-group-{gs} symmetric. " f"Activation: INT8 dynamic per-tensor scale (runtime, no calibration). " f"Fused kernel emitted by torch.compile via Inductor. " f"Skipped: wte, lm_head, norms, router, gate, small linears. " f"Verify val_loss vs fp16 baseline before a long run.", stacklevel=2, ) self.transformer.wte.weight = self.lm_head.weight assert self.transformer.wte.weight is self.lm_head.weight _init_device = ( torch.device("cuda", torch.cuda.current_device()) if torch.cuda.is_available() else torch.device("cpu") ) kernel.print_attention_backend(_init_device) kernel.warmup_attention_backend(_init_device) def set_tp_group(self, group) -> None: if self.config.tp_size <= 1: return for block in self.transformer.h: real_block = block.inner if isinstance(block, MoDBlock) else block real_block.attn.tp_group = group if hasattr(real_block, "ffn") and real_block.ffn is not None: real_block.ffn.tp_group = group if torch.distributed.is_initialized() and group is not None: for block in self.transformer.h: real_block = block.inner if isinstance(block, MoDBlock) else block attn = real_block.attn if not isinstance(attn, DenseMHA): continue if attn.num_kv_heads >= attn.tp_size: continue with torch.no_grad(): # Full qkv projection tensor is the canonical object to align. torch.distributed.broadcast(attn.qkv_proj.weight.data, src=0, group=group) if attn.qkv_proj.bias is not None: torch.distributed.broadcast(attn.qkv_proj.bias.data, src=0, group=group) def _init_weights(self, module): if isinstance(module, nn.Linear): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) if module.bias is not None: torch.nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): torch.nn.init.normal_(module.weight, mean=0.0, std=0.02) def get_num_params(self, non_embedding=True): """Count number of parameters""" n_params = sum(p.numel() for p in self.parameters()) if non_embedding: n_params -= self.transformer.wte.weight.numel() return n_params def forward(self, idx, targets=None, need_logits=True): idx = idx.long() if targets is not None: targets = targets.long() b, t = idx.size() torch._check( t <= self.config.block_size, lambda: f"Sequence too long: {t} > {self.config.block_size}", ) tok_emb = self.transformer.wte(idx) # (B, T, n_embd) x = self.transformer.drop(tok_emb) mod_gate_loss = torch.zeros((), device=x.device, dtype=torch.float32) moe_aux_loss = torch.zeros((), device=x.device, dtype=torch.float32) engram_hashes = self.engram_hasher(idx) if self.engram_hasher is not None else None for layer_i, (block, (is_mod, has_moe)) in enumerate(zip(self.transformer.h, self._block_flags)): if engram_hashes is not None and layer_i in self._engram_layer_ids: # Official placement: h = h + Engram(h, ids) BEFORE the block's attention. x = x + self.engram_layers[str(layer_i)](x, engram_hashes[layer_i]) if self.config.gradient_checkpointing and self.training: if is_mod: (x, moe_aux), gate_aux = torch.utils.checkpoint.checkpoint( block, x, use_reentrant=False ) mod_gate_loss = mod_gate_loss + gate_aux.float() elif has_moe: x, moe_aux = torch.utils.checkpoint.checkpoint( block, x, use_reentrant=False ) else: x = torch.utils.checkpoint.checkpoint( block, x, use_reentrant=False ) moe_aux = None elif is_mod: (x, moe_aux), gate_aux = block(x) mod_gate_loss = mod_gate_loss + gate_aux.float() elif has_moe: x, moe_aux = block(x) else: x = block(x) moe_aux = None if moe_aux is not None: moe_aux_loss = moe_aux_loss + moe_aux.float() pre_ln_f = x x = self.transformer.ln_f(x) if targets is not None: mtp_loss = torch.zeros((), device=x.device, dtype=torch.float32) if self.mtp is not None: mtp_loss = self.mtp.compute_loss( pre_ln_f, targets, self.lm_head.weight, mtp_lambda=self._mtp_lambda ) if not need_logits and _USE_LIGER and _LigerFusedLinearCrossEntropyLoss is not None and _liger_fused_ce_cuda_gate(x): liger_ce = _liger_fused_ce_attempt(x, self.lm_head.weight, targets) if liger_ce is not None: total_loss = liger_ce + mtp_loss + mod_gate_loss + moe_aux_loss self.last_ce_loss = liger_ce.detach() self.last_total_loss = total_loss.detach() return None, total_loss if not need_logits: ce_loss = chunked_cross_entropy(x, self.lm_head.weight, targets, ignore_index=-1) logits = None else: logits = F.linear(x, self.lm_head.weight) ce_loss = F.cross_entropy( logits.reshape(-1, logits.size(-1)).float(), targets.reshape(-1), ignore_index=-1, ) total_loss = ce_loss + mtp_loss + mod_gate_loss + moe_aux_loss self.last_ce_loss = ce_loss.detach() self.last_moe_aux_loss = moe_aux_loss.detach() self.last_total_loss = total_loss.detach() else: logits = F.linear(x[:, [-1], :], self.lm_head.weight) # (B, 1, vocab_size) ce_loss = None total_loss = None return logits, total_loss @torch.no_grad() def generate(self, idx, max_new_tokens=None, temperature=1.0, top_k=None): if max_new_tokens is None: max_new_tokens = getattr(self.config, 'max_gen_tokens', None) if max_new_tokens is None: raise ValueError( "generate() needs max_new_tokens, either passed " "directly or set as max_gen_tokens in config.py." ) if temperature is None or temperature <= 0: raise ValueError( f"generate(): temperature must be > 0, got {temperature!r}. " f"For greedy decoding use top_k=1 with a small temperature." ) # BUGFIX: @torch.no_grad() does NOT put the module in eval mode, so # sampling from a model left in .train() ran with dropout active (and # kept mutating MoE expert_bias). Force eval, restore afterwards. was_training = self.training self.eval() try: return self._generate_loop(idx, max_new_tokens, temperature, top_k) finally: if was_training: self.train() def _generate_loop(self, idx, max_new_tokens, temperature, top_k): for _ in range(max_new_tokens): idx_cond = idx if idx.size(1) <= self.config.block_size else idx[:, -self.config.block_size:] logits, _ = self(idx_cond) logits = logits[:, -1, :] / temperature if top_k is not None: v, _ = torch.topk(logits, min(top_k, logits.size(-1))) logits[logits < v[:, [-1]]] = -float('Inf') probs = F.softmax(logits.float(), dim=-1) # (B, vocab_size) # Sample next token idx_next = torch.multinomial(probs, num_samples=1) # (B, 1) # Append to sequence idx = torch.cat((idx, idx_next), dim=1) # (B, T+1) return idx def audit_fused_kernels(model, device, dtype=torch.float16, verbose=True): config = model.config head_dim = config.n_embd // config.n_head rope_dim = head_dim if getattr(config, 'd_rope', None) is None else config.d_rope b, t = 2, 8 results = {} if flash_ops is None: if verbose: print( "[FusionAudit] flash_ops extension failed to load entirely -- " "all three flash_* fusions (rope/outproj_add_rmsnorm/swiglu) " "are running eager fallback for the WHOLE run, not just under " "some condition. Check flash.cu compile output / nvcc " "availability above for the real cause." ) return {"flash_ops_loaded": False} results["flash_ops_loaded"] = True # -- fused RoPE -- try: q = torch.randn(b, t, config.n_head, head_dim, device=device, dtype=dtype) k = torch.randn(b, t, config.num_kv_heads or config.n_head, head_dim, device=device, dtype=dtype) cos, sin = build_rope_cache(config.block_size, rope_dim) cos, sin = cos.to(device), sin.to(device) out = flash_ops.fused_rope_qk(q, k, cos[:t].to(dtype), sin[:t].to(dtype), rope_dim) results["fused_rope"] = out is not None except Exception as e: # noqa: BLE001 -- diagnostic only, must not crash results["fused_rope"] = False results["fused_rope_error"] = repr(e) try: x = torch.randn(b * t, config.n_embd, device=device, dtype=dtype) w = torch.randn(config.n_embd, config.n_embd, device=device, dtype=dtype) residual = torch.randn(b * t, config.n_embd, device=device, dtype=dtype) norm_w = torch.ones(config.n_embd, device=device, dtype=dtype) out = flash_ops.fused_outproj_add_rmsnorm(x, w, residual, norm_w, 1e-5) results["fused_outproj_add_rmsnorm"] = out is not None except Exception as e: # noqa: BLE001 results["fused_outproj_add_rmsnorm"] = False results["fused_outproj_add_rmsnorm_error"] = repr(e) # -- fused SwiGLU -- try: hidden = int(config.n_embd * config.ffn_mult) gate = torch.randn(b * t, hidden, device=device, dtype=dtype) value = torch.randn(b * t, hidden, device=device, dtype=dtype) out = flash_ops.fused_swiglu(gate, value) results["fused_swiglu"] = out is not None except Exception as e: # noqa: BLE001 results["fused_swiglu"] = False results["fused_swiglu_error"] = repr(e) if verbose: print("[FusionAudit] one-time probe of fused kernels against real on-device tensors:") for name in ("fused_rope", "fused_outproj_add_rmsnorm", "fused_swiglu"): ok = results.get(name, False) status = "ENGAGED" if ok else "FELL BACK TO EAGER" print(f" {name:28s} -> {status}") if not ok and f"{name}_error" in results: print(f" reason: {results[f'{name}_error']}") n_ok = sum(results.get(n, False) for n in ("fused_rope", "fused_outproj_add_rmsnorm", "fused_swiglu")) if n_ok < 3: print( f" [FusionAudit] {3 - n_ok}/3 fusions NOT engaging -- given " f"this model is memory-bandwidth-bound (per the roofline " f"analysis), a missed fusion means real intermediate-tensor " f"HBM traffic that shouldn't be there. Worth fixing before " f"chasing anything else." ) else: print(" [FusionAudit] all 3/3 fusions engaged -- fusion is not the bottleneck here.") return results