Pragya / model.py
ArushBuilds's picture
Update model.py
4b04e69
Raw History Blame
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))
@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