Spaces:
Sleeping
Sleeping
Download model.py from ArushBuilds/Pragya: direct link, hf CLI and curl.
- Browser
- Download file 85.7 kB
-
https://huggingface.co/spaces/ArushBuilds/Pragya/resolve/61b358c953587a838c340fe97c699e643d70ce84/model.py
- Command line
-
hf download hf://spaces/ArushBuilds/Pragya@61b358c953587a838c340fe97c699e643d70ce84/model.py
-
curl -L -o model.py https://huggingface.co/spaces/ArushBuilds/Pragya/resolve/61b358c953587a838c340fe97c699e643d70ce84/model.py
85.7 kB
| 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)) | |
| 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) | |
| 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) | |
| 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 | |
| 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 | |
| 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) | |
| 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 | |
| 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 |