#!/usr/bin/env python3 """ NeoLLM model with FANformer, RMSNorm, ResFormer, Learnable Multipliers, full attention augmented with optional Momentum, MEA, and LUCID operators, Gated Attention (Qiu et al., 2025) combined with Affine-Scaled Attention (Bae et al., 2026), an optional Leviathan continuous token embedding generator (LEV layer, Batley & Saha 2026), optional Spelling Bee Embeddings (Rabe et al., 2026), optional Context Re-Positioning (Li et al., 2026), optional REPO-GRAPE contextual group positioning (Li et al., 2026 + Zhang et al., 2026), optional GOAT-style factorised attention log-priors (Litman & Guo, 2026), optional StackMemory (Zhang et al., NeurIPS 2025), and optional HOLA-derived episodic memory for the FlashAttention-2 IHA seq-expand path (Cui, 2026). When enabled, Leviathan-JTok/JTok-M uses reference PyTorch operations by default. The runtime-only ``config.jtok_kernel_backend = "triton"`` opt-in dispatches the same module parameters through the compatible CCE JTok/JTok-M adapter. The default ``"torch"`` path is unchanged and the backend choice does not add parameters or buffers to the checkpoint. The FlashAttention training paths deliberately use dense right-padded attention with a shared dtype and compilation contract across IHA and non-IHA. See ``NeoLLMModel.forward`` for its padding/packing constraints. Attention stack (orthogonal, all active simultaneously when enabled): 1. Gated Attention (use_gated_attention implicit via q_proj gate chunk): applies a head-specific elementwise sigmoid gate to the concatenated SDPA output before o_proj (G1 position, Qiu et al. 2025 §2.2). Introduces non-linearity between W_V and W_O, sparse input-dependent gating, and eliminates attention sink. 2. Affine-Scaled Attention (use_affine_scaled_attention): modulates softmax attention weights directly as [α(X)·softmax(QK^T/√dk) + β(X)] V relaxing the unit-sum constraint of softmax. α is per-head, per-query, input-dependent and bounded in [0,1] via linear_clipping. β is a moving-average bias that prevents collapse. Reduces first-token bias, increases attention entropy, and is complementary to Gated Attention (Bae et al. 2026, Table 2: Affine-Scaled + Gated > either alone). Flash/SDPA path: Expanding [α·softmax(QKᵀ)+β]V distributively yields two terms: α · attention(Q,K,V) — backend computes this directly + β · Σ_{j∈A(i)} V_j — sum over the same valid keys A(i) For global causal attention A(i) is a prefix and this is V.cumsum. For sliding-window or packed sequences the sum is windowed or segmented so β cannot leak outside the backend attention mask. Per-weight tensors (attn_weights_pre/post_affine) are unavailable in flash mode. Eager path: applies the same valid-key mask to α·softmax + β and keeps full weight access for interpretability. References: FANformer: "FANformer: Improving Large Language Models Through Effective Periodicity Modeling" Learnable Multipliers: "Learnable Multipliers: Freeing the Scale of Language Model Matrix Layers" Gated Attention: Qiu et al. (2025). "Gated Attention for Large Language Models: Non-linearity, Sparsity, and Attention-Sink-Free." arXiv:2505.06708. Affine-Scaled Attention: Bae et al. (2026). "Affine-Scaled Attention: Towards Flexible and Stable Transformer Attention." arXiv:2602.23057. Leviathan Generator: Batley & Saha (2026). "A Separable Architecture for Continuous Token Representation in Language Models." arXiv:2601.22040. JTok: Yang et al. (2026). "JTok: On Token Embedding as another Axis of Scaling Law via Joint Token Self-modulation." arXiv:2602.00800. KHRONOS: Batley & Saha (2025). "KHRONOS: a Kernel-Based Neural Architecture for Rapid, Resource-Efficient Scientific Computation." arXiv:2505.13315. Spelling Bee Embeddings: Rabe, Clymo & Dong (2026). "Spelling Bee Embeddings for Language Modeling." arXiv:2601.18030. Context Re-Positioning: Li, Zhao, Cai & Sproat (2026). "REPO: Language Models with Context Re-Positioning." arXiv:2512.14391. GRAPE: Zhang et al. (2026). "Group Representational Position Encoding." arXiv:2512.07805 / ICLR 2026. GOAT priors: Litman & Guo (2026). "You Need Better Attention Priors." arXiv:2601.15380. StackMemory / STACKTRANS: Zhang et al. (2025). "Recursive Transformer: Boosting Reasoning Ability with State Stack." NeurIPS 2025. HOLA: Cui (2026). "A Hippocampus for Linear Attention: An Exact Memory for What the Recurrent State Forgets." arXiv:2607.02303. TWEO: Liang et al. (2026). "Transformers Without Extreme Outliers Enables FP8 Training And Quantization For Dummies." arXiv:2511.23225. """ import inspect import math from dataclasses import dataclass from typing import Optional, Union, Tuple import torch import torch.nn.functional as F from torch import nn from torch.utils.checkpoint import checkpoint as _activation_checkpoint try: from cut_cross_entropy import linear_cross_entropy # type: ignore[import-not-found] _CCE_AVAILABLE = True except ImportError: linear_cross_entropy = None _CCE_AVAILABLE = False try: from cut_cross_entropy import meap_mask_inputs # type: ignore[import-not-found] except ImportError: meap_mask_inputs = None try: from cut_cross_entropy.polynorm import ( # type: ignore[import-not-found] polynorm as _cce_polynorm, ) except (ImportError, ModuleNotFoundError): _cce_polynorm = None try: from cut_cross_entropy.polynorm import ( # type: ignore[import-not-found] polynorm_uses_cute as _cce_polynorm_uses_cute, ) except (ImportError, ModuleNotFoundError): _cce_polynorm_uses_cute = None _REPO_GRAPE_KERNEL_AVAILABLE = False _cce_repo_grape = None _cce_repo_grape_supported = None try: from cut_cross_entropy.repo_grape import ( # type: ignore[import-not-found] repo_grape as _cce_repo_grape, repo_grape_supported as _cce_repo_grape_supported, ) _REPO_GRAPE_KERNEL_AVAILABLE = True except (ImportError, ModuleNotFoundError, AttributeError): # REPO-GRAPE remains fully functional through RepoGrapePositioning when the # optional Triton package (or a new-enough CCE checkout) is unavailable. pass _LEV_KERNEL_AVAILABLE = False _leviathan_embedding_compiler_safe = None _leviathan_embedding_with_seed_compiler_safe = None _leviathan_supports = None try: from cut_cross_entropy.leviathan import ( # type: ignore[import-not-found] _HAS_KERNEL_FORWARD as _LEV_HAS_KERNEL_FORWARD, leviathan_embedding_compiler_safe as _leviathan_embedding_compiler_safe, supports as _leviathan_supports, ) _LEV_KERNEL_AVAILABLE = bool(_LEV_HAS_KERNEL_FORWARD) try: from cut_cross_entropy.leviathan import ( # type: ignore[import-not-found] leviathan_embedding_with_seed_compiler_safe as _leviathan_embedding_with_seed_compiler_safe, ) except (ImportError, ModuleNotFoundError, AttributeError): # An older CCE package may provide the legacy Leviathan kernel without # the opt-in JTok geometry bridge. Keep the legacy path available. pass except (ImportError, ModuleNotFoundError, AttributeError): # The model must remain usable with a plain CCE installation. In that # environment the original reference LeviathanGenerator.forward below is # the complete embedding path. _LEV_HAS_KERNEL_FORWARD = False _LEV_GEOMETRY_KERNEL_AVAILABLE = ( _LEV_KERNEL_AVAILABLE and _leviathan_embedding_with_seed_compiler_safe is not None ) # JTok/JTok-M is a separate, strict opt-in path. Keep this import separate # from the legacy Leviathan import above: an older CCE installation may expose # the original Leviathan kernel but not the newer JTok adapter, and that must # not silently disable the legacy path. _JTOK_KERNEL_AVAILABLE = False _apply_neollm_jtok = None try: from cut_cross_entropy.leviathan import ( # type: ignore[import-not-found] apply_neollm_jtok as _apply_neollm_jtok, ) _JTOK_KERNEL_AVAILABLE = True except (ImportError, ModuleNotFoundError, AttributeError): pass def _leviathan_config_supported(config) -> bool: """Check LEV shapes without depending on NeoLLM's optional dtype field.""" if _leviathan_supports is None: return False try: # NeoLLM's PretrainedConfig.dtype may be None even after # ``module.to(torch.bfloat16)``. The parameter checks below are the # source of truth; the override keeps this predicate trace-friendly. return bool(_leviathan_supports(config, dtype_override=torch.bfloat16)) except (AttributeError, TypeError, ValueError): return False try: from liger_kernel.transformers import ( # type: ignore[import-not-found] LigerFusedLinearCrossEntropyLoss, ) _LIGER_KERNEL_AVAILABLE = True except ImportError: LigerFusedLinearCrossEntropyLoss = None _LIGER_KERNEL_AVAILABLE = False from transformers.generation import GenerationMixin from transformers.masking_utils import create_causal_mask from transformers.modeling_flash_attention_utils import ( FlashAttentionKwargs, _flash_attention_forward, flash_attn_supports_top_left_mask, ) from transformers.modeling_layers import GradientCheckpointingLayer from transformers.modeling_outputs import ( BaseModelOutputWithPast, CausalLMOutputWithPast, ) from transformers.modeling_rope_utils import ROPE_INIT_FUNCTIONS, dynamic_rope_update from transformers.modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel from transformers.processing_utils import Unpack from transformers.utils import TransformersKwargs, logging from .configuration_neollm import NeoLLMConfig from transformers import AutoConfig, AutoModel, AutoModelForCausalLM @dataclass class NeoLLMBaseModelOutputWithPast(BaseModelOutputWithPast): """Backbone output with compact, differentiable JTok-M statistics. When JTok-M is active during training, ``jtokm_aux_stats`` contains one ``(P, n, T)`` tuple per decoder layer. ``P`` is the differentiable sum of the router probabilities, ``n`` is the detached hard Top-K count, and ``T`` is the number of valid routed tokens. Keeping only these reductions in the model output avoids retaining per-token router tensors until the language-model loss is assembled. """ jtokm_aux_stats: Optional[ Tuple[Tuple[torch.Tensor, torch.Tensor, torch.Tensor], ...] ] = None # StackMemory diagnostics packed as [num_layers, 11]. Column 0 retains # the differentiable action entropy used by the auxiliary StackTrans loss; # the remaining columns are detached monitoring reductions. stack_metrics: Optional[torch.Tensor] = None logger = logging.get_logger(__name__) _NEOLLM_FA_USE_TOP_LEFT_MASK = flash_attn_supports_top_left_mask() # ── Optional flash_attn direct import (required only for IHA seq-expand, P>1) ── # Attempted once at module load time so the symbols are available when # NeoLLMAttention.__init__ stores them as instance attributes. If flash_attn # is not installed the names are set to None and a clear ImportError is raised # inside __init__ only when the user actually enables IHA with P>1. # The common path (P=1 or use_iha=False) never triggers the error. # # Functions surfaced: # flash_attn_func – standard batched causal attention (no padding path) # flash_attn_varlen_func – variable-length / packed-sequence causal attention # # Reference: Dao-AILab/flash-attention, flash_attn.flash_attn_interface try: from flash_attn.flash_attn_interface import ( # type: ignore[import] flash_attn_func as _IHA_FA_FUNC, flash_attn_varlen_func as _IHA_FA_VARLEN, ) _IHA_FLASH_ATTN_AVAILABLE = True except ImportError: _IHA_FA_FUNC = None _IHA_FA_VARLEN = None _IHA_FLASH_ATTN_AVAILABLE = False class ScalarMultiplier(nn.Module): """ Scalar Learnable Multiplier: W̃ = s·W From "Learnable Multipliers: Freeing the Scale of Language Model Matrix Layers": Allows the effective matrix norm ||W̃|| = s·||W|| to adapt to data, escaping the WD-noise equilibrium that constrains ||W|| ∝ √(η/λ). """ def __init__(self, initial_value: float = 1.0): super().__init__() self.multiplier = nn.Parameter(torch.tensor(initial_value)) def forward(self, x: torch.Tensor) -> torch.Tensor: return self.multiplier * x class VectorMultiplier(nn.Module): """ Vector Learnable Multipliers: W̃ = diag(r)·W·diag(c) From "Learnable Multipliers: Freeing the Scale of Language Model Matrix Layers": Frees not only the overall matrix norm but also individual row/column norms from the WD-noise equilibrium, enabling richer feature scale diversity. """ def __init__( self, dim: int, multiplier_type: str = "row", initial_value: float = 1.0 ): super().__init__() self.multiplier_type = multiplier_type self.multiplier = nn.Parameter(torch.ones(dim) * initial_value) def forward(self, x: torch.Tensor) -> torch.Tensor: return x * self.multiplier class LinearWithMultipliers(nn.Module): """ Linear layer with optional row and/or column learnable multipliers. Implements: y = (r ⊙ (W @ (c ⊙ x))) + b when enabled. With ``enable_multipliers=False`` no multiplier parameters are instantiated and the module reduces to its wrapped ``nn.Linear``. """ def __init__( self, in_features: int, out_features: int, bias: bool = True, use_row_multiplier: bool = False, use_column_multiplier: bool = False, enable_multipliers: bool = True, ): super().__init__() self.linear = nn.Linear(in_features, out_features, bias=bias) self.enable_multipliers = bool(enable_multipliers) self.use_row_multiplier = bool(use_row_multiplier and self.enable_multipliers) self.use_column_multiplier = bool( use_column_multiplier and self.enable_multipliers ) if self.use_row_multiplier: self.row_multiplier = VectorMultiplier(out_features, multiplier_type="row") if self.use_column_multiplier: self.column_multiplier = VectorMultiplier( in_features, multiplier_type="column" ) def forward(self, x: torch.Tensor) -> torch.Tensor: if self.use_column_multiplier: x = self.column_multiplier(x) x = self.linear(x) if self.use_row_multiplier: x = self.row_multiplier(x) return x class EmbeddingWithMultipliers(nn.Module): """ Token embedding matrix with optional vocabulary-row and hidden-channel learnable multipliers. Effective embedding: E_eff[token, channel] = r_token * E[token, channel] * c_channel This mirrors the Learnable Multipliers paper's embedding recommendation, while keeping the actual weight available as ``.weight`` for tooling that expects a standard ``nn.Embedding``-like object. It is only intended for the untied, non-generator embedding path. """ def __init__( self, num_embeddings: int, embedding_dim: int, padding_idx: Optional[int] = None, enable_multipliers: bool = True, ): super().__init__() self.embedding = nn.Embedding( num_embeddings, embedding_dim, padding_idx=padding_idx ) self.num_embeddings = num_embeddings self.embedding_dim = embedding_dim self.padding_idx = padding_idx self.enable_multipliers = bool(enable_multipliers) self.use_row_multiplier = self.enable_multipliers self.use_column_multiplier = self.enable_multipliers if self.enable_multipliers: self.row_multiplier = nn.Parameter(torch.ones(num_embeddings)) self.column_multiplier = VectorMultiplier( embedding_dim, multiplier_type="column" ) @property def weight(self) -> torch.nn.Parameter: return self.embedding.weight def forward(self, input_ids: torch.Tensor) -> torch.Tensor: x = self.embedding(input_ids) if self.use_column_multiplier: x = self.column_multiplier(x) if self.use_row_multiplier: row_scale = F.embedding(input_ids, self.row_multiplier.unsqueeze(-1)).to( dtype=x.dtype ) x = x * row_scale return x # ==================== LEVIATHAN CONTINUOUS TOKEN GENERATOR ==================== # Sampled dynamics: mathematical definitions and interpretation # ---------------------------------------------------------- # Enable through USE_DYNAMICS_METRICS in train.py before model construction. # train.py supplies config.dynamics_metrics_enabled and dynamics_sample_tokens; # a standalone caller may set those same attributes before constructing a model. # Collection requires native CUDA/Triton checkpoints, training mode and enabled # autograd. It is detached, read-only, and never contributes to the objective. # # Sampling/aggregation contract: # Flatten batch and sequence to N rows, choose S=min(N, sample_tokens) positions # r_j=floor(j*N/S), then exclude invalid sampled rows. There is no resampling # to replace masked rows. Explicit attention validity excludes padding; MEAP # replacement rows are excluded from generator diagnostics. Token IDs alone # are not a padding test (EOS may share the padding ID). JToK uses its existing # validity mask. Let V be the remaining sampled rows and A(f)=mean_{r in V} f_r. # Every row statistic below is reduced with A, except sample_valid_rows=|V|. # Thus A(RMS(v_r)) is NOT RMS over all flattened activations. Empty V produces # neutral zeros; consult sample_valid_rows before interpreting any metric. # Deterministic position sampling is not an unbiased dataset estimate and can # alias positional structure. The cache describes the last gradient-enabled # microbatch/recomputation, not an optimizer-step, epoch or validation average. # Arithmetic uses FP32 on detached, actually stored native tensors. Nonfinite # values are not silently repaired; interpret RMS/angles with finite fractions. # # Leviathan metrics (prefix dynamics/leviathan/): # Let e_r be the generated embedding, z_r the seed, and m_r the concatenation # of stored CP product modes across heads/ranks. Define RMS(v)=sqrt(mean(v^2)), # Z(v)=mean(1[v=0]), NF(v)=mean(1[not finite(v)]). # embedding_rms = A(RMS(e)); tracks generator output scale. # seed_rms = A(RMS(z)); tracks the shared seed representation's scale. # product_mode_rms = A(RMS(m)); tracks the CP branch's numerical magnitude. # embedding_nonfinite_fraction = A(NF(e)); flags invalid generated outputs. # product_mode_zero_fraction = A(Z(m)); flags stored zero/underflow/collapse, # without distinguishing their causes or the separate residual branch. # product_mode_nonfinite_fraction = A(NF(m)); localizes invalid CP products. # coordinate_saturation_fraction = A(mean_{head,d} 1[c<=.01 or c>=.99]), # c=sigmoid(.5*(gamma*xhat+beta)). xhat/gamma/beta are native checkpoints; # c is reconstructed only for sampled rows in FP32, before any activation # cast. This is a boundary/conditioning proxy, not the exact stored # coordinate's threshold test or proof that spline derivatives vanish. # Seed, mode and embedding scales separate representation growth from product # attenuation and projection effects. Their trends alone do not prove learning. # # JToK versus JToK-M: keep their distinct forward semantics when comparing. # x is the incoming MLP update delta_m, not the full decoder hidden state; # y is the update returned by this module, d=y-x, s is the learned scaler. # With n(v)=v/(||v||_2+norm_eps), on valid rows: # JToK: y=x*(1+s*n(S(z_tilde))) [elementwise modulation]. # JToK-M: y=x+a*s*n(sum_k p_k*S_{e_k}(z_tilde)) [additive residual], # a=1/sqrt(2L), e_k=TopK(router_logits), # p_k=sigmoid(logit_{e_k})/max(sum_j sigmoid(logit_{e_j}),norm_eps). # Native stored output rounding is part of the observed d. The observer does # not reconstruct an ideal gate/residual or change normalization/routing. # # Shared metrics (prefix dynamics/jtok|jtokm/layer_/), eps=1e-12: # input_rms=A(RMS(x)), output_rms=A(RMS(y)), update_rms=A(RMS(d)); # separate incoming scale, outgoing scale and effective intervention. # relative_update_norm=A(||d||_2/max(||x||_2,eps)); unbounded effect size, # sensitive to near-zero input. Read alongside zero_input_fraction. # usage_index=A(||d||_2/max(||x||_2+||d||_2,eps)); bounded [0,1] for finite # inputs, zero for an identity intervention. It measures visible effect, # NOT probability of use, benefit, gradient flow or semantic relevance. # input_output_cosine=A(clip(/max(||x||_2*||y||_2,eps),-1,1)); # detects direction preservation/reversal separately from scale changes. # update_alignment=A(clip(/max(||x||_2*||d||_2,eps),-1,1)); # positive/negative values indicate reinforcement/cancellation of x. # orthogonal_update_fraction=A(1-alignment_r^2 if ||d||_2>0 else 0); # for nonzero, well-scaled x this is the fraction of d's energy # orthogonal to x. With x=0,d!=0 the convention is 1 (no reference # direction); epsilon-dominated angles are proxies, not exact geometry. # It is the mean of per-row fractions, NOT 1-A(alignment)^2. # changed_element_fraction=A(mean_h 1[y_h!=x_h]); visible dtype-level # intervention, not an FP32-master update or usefulness test. # zero_input_fraction=A(1[||x||_2=0]); qualifies norm ratios and angles. # output_nonfinite_fraction=A(NF(y)); output numerical-health check. # product_mode_zero_fraction=A(Z(m)), product_mode_nonfinite_fraction=A(NF(m)); # m contains all stored selected-expert CP modes (one expert for JToK). # coordinate_saturation_fraction=A(mean_d 1[z_tilde_d<=.01 or >=.99]); # observes the supplied coordinate, not the generator head coordinates. # Jointly usage, alignment and orthogonal fraction distinguish a small/identity # intervention, reinforcement/cancellation and direction-changing corrections. # They cannot establish why a token needs that correction; causal usefulness # requires a controlled ablation and task/loss evidence, not these proxies. # # Additional JToK-M routing metrics (not emitted for plain JToK): # For selected stored weights p, b=sum_k p_k and pi_k=p_k/b when b>0 (else 0): # selected_weight_entropy=A(-sum_k pi_k*log(pi_k)); entropy in nats over the # selected K, not all E experts. K=1 is necessarily zero. Diagnostic pi # renormalization does not modify the actual p used by the model. # selected_max_weight=A(max_k p_k); concentration of actual mixture weights. # selected_weight_sum=A(b), zero_selected_mass_fraction=A(1[b=0]); reveal # epsilon-limited or underflowed mixture mass; entropy alone can hide it. # Let f_e=sum_{r in V,k}1[e_k(r)=e]/(|V|*K), selection counts, NOT weight mass: # expert__selection_fraction=f_e; sampled expert traffic. # load_cv=std_population(f)/max(mean(f),eps); imbalance across experts. # load_entropy_normalized=-sum_e f_e*log(max(f_e,1e-30))/max(log(E),eps); # coverage/uniformity on [0,1] when V is nonempty and E>1; E=1 yields 0. # active_experts=sum_e 1[f_e>0]; observed coverage, not causal contribution. # These characterize routing collapse/specialization candidates, not expert # quality. Low load entropy may be intentional specialization. Always compare # with usage, update geometry, numerical health and train/validation loss. # # Cost/API contract: each observed forward samples at most sample_tokens rows, # reduces on GPU and caches only compact detached FP32 scalars, never token IDs # or activation vectors. There are extra GPU launches/workspaces, so cheap does # not mean free. get_dynamics_metrics() exposes device scalars outside compiled # forward; train.py batches their host transfer only when logging. Evaluation # does not collect these training snapshots or relabel them as eval metrics. try: import triton import triton.language as tl except (ImportError, ModuleNotFoundError): triton = tl = None LEV_NAMES = ( "sample_valid_rows", "embedding_rms", "embedding_nonfinite_fraction", "seed_rms", "coordinate_saturation_fraction", "product_mode_rms", "product_mode_zero_fraction", "product_mode_nonfinite_fraction", ) JTOK_NAMES = ( "sample_valid_rows", "input_rms", "output_rms", "update_rms", "relative_update_norm", "input_output_cosine", "changed_element_fraction", "product_mode_zero_fraction", "product_mode_nonfinite_fraction", "coordinate_saturation_fraction", "selected_weight_entropy", "selected_max_weight", "output_nonfinite_fraction", "zero_input_fraction", "usage_index", "update_alignment", "orthogonal_update_fraction", "selected_weight_sum", "zero_selected_mass_fraction", ) if triton is not None: @triton.jit def _dynamics_reduce_kernel(rows, output, S: tl.constexpr, C: tl.constexpr, BS: tl.constexpr): col = tl.program_id(0) r = tl.arange(0, BS) valid = tl.load(rows + r * C, r < S, 0) count = tl.sum(valid, 0) values = tl.load(rows + r * C + col, r < S, 0) value = tl.sum(tl.where(valid > 0, values, 0), 0) / tl.maximum(count, 1) tl.store(output + col, tl.where(col == 0, count, value)) @triton.jit def _jtok_dynamics_rows_kernel(delta, out, z, modes, experts, weights, valid, rows, N: tl.constexpr, S: tl.constexpr, H: tl.constexpr, D: tl.constexpr, M: tl.constexpr, K: tl.constexpr, E: tl.constexpr, BH: tl.constexpr, BD: tl.constexpr, BM: tl.constexpr, BK: tl.constexpr, HAS_MASK: tl.constexpr): sample = tl.program_id(0) row = sample * N // S is_valid = tl.full((), True, tl.int1) if HAS_MASK: is_valid = tl.load(valid + row) h = tl.arange(0, BH) x = tl.load(delta + row * H + h, h < H, 0).to(tl.float32) y = tl.load(out + row * H + h, h < H, 0).to(tl.float32) diff = y - x xn = tl.sqrt(tl.sum(x * x, 0)) yn = tl.sqrt(tl.sum(y * y, 0)) dn = tl.sqrt(tl.sum(diff * diff, 0)) cosine = tl.minimum(1., tl.maximum(-1., tl.sum(x * y, 0) / tl.maximum(xn * yn, 1e-12))) d = tl.arange(0, BD) coord = tl.load(z + row * D + d, d < D, 0).to(tl.float32) saturation = tl.sum(((coord <= .01) | (coord >= .99)) & (d < D), 0) / D m = tl.arange(0, BM) mode = tl.load(modes + row * K * M + m, m < K * M, 0).to(tl.float32) mz = tl.sum((mode == 0) & (m < K * M), 0) / (K * M) mn = tl.sum(((mode != mode) | (tl.abs(mode) == float('inf'))) & (m < K * M), 0) / (K * M) k = tl.arange(0, BK) p = tl.load(weights + row * K + k, k < K, 0).to(tl.float32) mass = tl.sum(p, 0) probability = p / tl.where(mass > 0, mass, 1.) entropy = -tl.sum(tl.where(probability > 0, probability * tl.log(tl.maximum(probability, 1e-30)), 0), 0) max_weight = tl.max(p, 0) nonfinite = tl.sum(((y != y) | (tl.abs(y) == float('inf'))) & (h < H), 0) / H changed = tl.sum((y != x) & (h < H), 0) / H alignment = tl.minimum(1., tl.maximum(-1., tl.sum(x * diff, 0) / tl.maximum(xn * dn, 1e-12))) orthogonal = tl.where(dn > 0, tl.maximum(0., 1. - alignment * alignment), 0.) C: tl.constexpr = 19 + E base = rows + sample * C tl.store(base, is_valid.to(tl.float32)) tl.store(base + 1, xn / tl.sqrt(float(H))) tl.store(base + 2, yn / tl.sqrt(float(H))) tl.store(base + 3, dn / tl.sqrt(float(H))) tl.store(base + 4, dn / tl.maximum(xn, 1e-12)) tl.store(base + 5, cosine) tl.store(base + 6, changed) tl.store(base + 7, mz) tl.store(base + 8, mn) tl.store(base + 9, saturation) tl.store(base + 10, entropy) tl.store(base + 11, max_weight) tl.store(base + 12, nonfinite) tl.store(base + 13, (xn == 0).to(tl.float32)) tl.store(base + 14, dn / tl.maximum(xn + dn, 1e-12)) tl.store(base + 15, alignment) tl.store(base + 16, orthogonal) tl.store(base + 17, mass) tl.store(base + 18, (mass == 0).to(tl.float32)) route = tl.load(experts + row * K + k, k < K, -1) for expert in range(E): fraction = tl.sum((route == expert) & (k < K), 0) / K tl.store(base + 19 + expert, fraction) @triton.jit def _lev_dynamics_rows_kernel(embedding, seed, xhat, norm_weight, norm_bias, modes, valid, rows, N: tl.constexpr, S: tl.constexpr, H: tl.constexpr, D: tl.constexpr, HEADS: tl.constexpr, R: tl.constexpr, BH: tl.constexpr, BD: tl.constexpr, BM: tl.constexpr, HAS_MASK: tl.constexpr): sample = tl.program_id(0) row = sample * N // S is_valid = tl.full((), True, tl.int1) if HAS_MASK: is_valid = tl.load(valid + row) h = tl.arange(0, BH) value = tl.load(embedding + row * H + h, h < H, 0).to(tl.float32) erms = tl.sqrt(tl.sum(value * value, 0) / H) en = tl.sum(((value != value) | (tl.abs(value) == float('inf'))) & (h < H), 0) / H d = tl.arange(0, BD) z = tl.load(seed + row * D + d, d < D, 0).to(tl.float32) zrms = tl.sqrt(tl.sum(z * z, 0) / D) saturated = tl.full((), 0., tl.float32) for head in range(HEADS): x = tl.load(xhat + (head * N + row) * D + d, d < D, 0).to(tl.float32) w = tl.load(norm_weight + head * D + d, d < D, 0).to(tl.float32) b = tl.load(norm_bias + head * D + d, d < D, 0).to(tl.float32) coordinate = tl.sigmoid(.5 * (x * w + b)) saturated += tl.sum(((coordinate <= .01) | (coordinate >= .99)) & (d < D), 0) m = tl.arange(0, BM) mode = tl.load(modes + row * HEADS * R + m, m < HEADS * R, 0).to(tl.float32) mrms = tl.sqrt(tl.sum(mode * mode, 0) / (HEADS * R)) mz = tl.sum((mode == 0) & (m < HEADS * R), 0) / (HEADS * R) mn = tl.sum(((mode != mode) | (tl.abs(mode) == float('inf'))) & (m < HEADS * R), 0) / (HEADS * R) base = rows + sample * 8 tl.store(base, is_valid.to(tl.float32)) tl.store(base + 1, erms) tl.store(base + 2, en) tl.store(base + 3, zrms) tl.store(base + 4, saturated / (HEADS * D)) tl.store(base + 5, mrms) tl.store(base + 6, mz) tl.store(base + 7, mn) def _check_sampling(max_samples: int) -> None: validate_dynamics_samples(max_samples) if max_samples == 0: raise ValueError("diagnostics max_samples must be in [1, 4096]") def validate_dynamics_samples(samples: int) -> None: """Zero disables collection; reject invalid budgets before native work.""" if type(samples) is not int or not 0 <= samples <= 4096: raise ValueError("dynamics_samples must be an integer in [0, 4096]") def _check_tensors(primary: torch.Tensor, *others: torch.Tensor) -> None: for value in (primary, *others): if value.device != primary.device or not value.is_contiguous(): raise ValueError("diagnostics tensors must be contiguous on the same device") if value.dtype not in (torch.float16, torch.bfloat16, torch.float32): raise TypeError("diagnostics floating tensors require FP16, BF16 or FP32") def _check_mask(valid: torch.Tensor, n: int, device: torch.device) -> None: if (valid.dtype != torch.bool or valid.device != device or valid.ndim != 1 or valid.numel() not in (0, n) or not valid.is_contiguous()): raise ValueError("diagnostics mask must be a contiguous bool vector of length zero or N on the input device") def _reduce(rows: torch.Tensor, columns: int) -> torch.Tensor: output = torch.empty(columns, device=rows.device, dtype=torch.float32) _dynamics_reduce_kernel[(columns,)]( rows, output, rows.shape[0], columns, triton.next_power_of_2(rows.shape[0]), num_warps=4) return output @torch.library.custom_op("neollm_dynamics::jtok_dynamics", mutates_args=()) def jtok_dynamics(delta: torch.Tensor, output: torch.Tensor, z: torch.Tensor, modes: torch.Tensor, experts: torch.Tensor, weights: torch.Tensor, valid: torch.Tensor, num_experts: int, max_samples: int) -> torch.Tensor: _check_sampling(max_samples) if type(num_experts) is not int or not 1 <= num_experts <= 1024: raise ValueError("diagnostics num_experts must be an integer in [1, 1024]") if delta.ndim != 2 or delta.shape[1] == 0 or output.shape != delta.shape: raise ValueError("diagnostics input/output must have shape [N, H] with H > 0") n, hidden = delta.shape if z.ndim != 2 or z.shape[0] != n or z.shape[1] == 0: raise ValueError("diagnostics coordinates must have shape [N, D] with D > 0") if (modes.ndim != 3 or modes.shape[0] != n or modes.shape[2] == 0 or not 1 <= modes.shape[1] <= num_experts): raise ValueError("diagnostics modes must have shape [N, K, M] with 1 <= K <= E and M > 0") if experts.shape != modes.shape[:2] or weights.shape != experts.shape: raise ValueError("diagnostics routes/weights must have shape [N, K]") if (experts.dtype not in (torch.int32, torch.int64) or experts.device != delta.device or not experts.is_contiguous()): raise ValueError("diagnostics routes must be contiguous integer tensors on the input device") _check_tensors(delta, output, z, modes, weights) _check_mask(valid, n, delta.device) if triton is None or not delta.is_cuda: raise RuntimeError("native JToK diagnostics require CUDA/Triton") samples = min(n, max_samples) columns = len(JTOK_NAMES) + num_experts if not samples: return torch.zeros(columns, device=delta.device, dtype=torch.float32) rows = torch.empty((samples, columns), device=delta.device, dtype=torch.float32) _jtok_dynamics_rows_kernel[(samples,)]( delta, output, z, modes, experts, weights, valid, rows, n, samples, hidden, z.shape[1], modes.shape[2], modes.shape[1], num_experts, triton.next_power_of_2(hidden), triton.next_power_of_2(z.shape[1]), triton.next_power_of_2(modes.shape[1] * modes.shape[2]), triton.next_power_of_2(modes.shape[1]), valid.numel() != 0, num_warps=4) return _reduce(rows, columns) @jtok_dynamics.register_fake def _jtok_dynamics_fake(delta, output, z, modes, experts, weights, valid, num_experts, max_samples): return torch.empty(len(JTOK_NAMES) + num_experts, device=delta.device, dtype=torch.float32) @torch.library.custom_op("neollm_dynamics::leviathan_dynamics", mutates_args=()) def leviathan_dynamics(embedding: torch.Tensor, seed: torch.Tensor, xhat: torch.Tensor, norm_weight: torch.Tensor, norm_bias: torch.Tensor, modes: torch.Tensor, valid: torch.Tensor, max_samples: int) -> torch.Tensor: _check_sampling(max_samples) if embedding.ndim < 2 or embedding.shape[-1] == 0: raise ValueError("diagnostics embedding must have shape [..., H] with H > 0") n, hidden = embedding.numel() // embedding.shape[-1], embedding.shape[-1] if (xhat.ndim != 3 or xhat.shape[0] == 0 or xhat.shape[1] != n or xhat.shape[2] == 0): raise ValueError("diagnostics xhat must have shape [heads, N, D] with positive heads and D") heads, _, d = xhat.shape if seed.shape != (n, d) or norm_weight.shape != (heads, d) or norm_bias.shape != (heads, d): raise ValueError("diagnostics seed/norm metadata do not match xhat") if modes.ndim != 3 or modes.shape[:2] != (n, heads) or modes.shape[2] == 0: raise ValueError("diagnostics modes must have shape [N, heads, rank] with rank > 0") _check_tensors(embedding, seed, xhat, norm_weight, norm_bias, modes) _check_mask(valid, n, embedding.device) if triton is None or not embedding.is_cuda: raise RuntimeError("native Leviathan diagnostics require CUDA/Triton") samples = min(n, max_samples) if not samples: return torch.zeros(len(LEV_NAMES), device=embedding.device, dtype=torch.float32) rows = torch.empty((samples, len(LEV_NAMES)), device=embedding.device, dtype=torch.float32) rank = modes.shape[-1] _lev_dynamics_rows_kernel[(samples,)]( embedding, seed, xhat, norm_weight, norm_bias, modes, valid, rows, n, samples, hidden, d, heads, rank, triton.next_power_of_2(hidden), triton.next_power_of_2(d), triton.next_power_of_2(heads * rank), valid.numel() != 0, num_warps=4) return _reduce(rows, len(LEV_NAMES)) @leviathan_dynamics.register_fake def _leviathan_dynamics_fake(embedding, seed, xhat, norm_weight, norm_bias, modes, valid, max_samples): return torch.empty(len(LEV_NAMES), device=embedding.device, dtype=torch.float32) def unpack_jtok_dynamics(packed: torch.Tensor, num_experts: int, *, mixture: bool): """Expose device scalars; the caller chooses when to transfer to its logger.""" result = dict(zip(JTOK_NAMES, packed[:len(JTOK_NAMES)].unbind())) if not mixture: result.pop("selected_weight_entropy") result.pop("selected_max_weight") result.pop("selected_weight_sum") result.pop("zero_selected_mass_fraction") return result fractions = packed[len(JTOK_NAMES):] mean = fractions.mean() result["load_cv"] = fractions.std(unbiased=False) / mean.clamp_min(1e-12) result["load_entropy_normalized"] = ( -(fractions * fractions.clamp_min(1e-30).log()).sum() / torch.tensor(float(num_experts), device=packed.device).log().clamp_min(1e-12)) result["active_experts"] = (fractions > 0).sum().float() for expert, value in enumerate(fractions.unbind()): result[f"expert_{expert}_selection_fraction"] = value return result class LeviathanGenerator(nn.Module): """ Learned Embedding Vectorization (LEV) layer - Leviathan input embedding (Batley & Saha, 2026, Sec. 3.1; arXiv:2601.22040). Replaces the dense input embedding matrix E in R^{VxD} with a compact continuous generator G : {0,...,V-1} -> R^D: 1. Compositional indexing (base-b). The vocabulary is factorized into k components via a base-b decomposition, i -> (i_1, ..., i_k), with b = ceil(V^(1/k)). Each component indexes a shared codebook C_r in R^{b x d_seed} and the seed is z(i) = sum_{r=1..k} C_r[i_r]. 2. Per-head seed projection. For each of h heads l the seed is projected and normalized to a bounded latent coordinate z~_l = sigmoid(1/2 * LN(W_seed,l z(i))) in [0,1]^{d_seed}. The 1/2 scaling factor keeps sigmoid inputs near the linear regime. 3. Separable B-spline expansion. Each dimension of z~_l is expanded on a univariate quadratic B-spline basis with kappa knots over [0,1], and the per-dimension expansions are combined by a rank-r tensor product with sign-parity tracking: phi_l[d,k] = sum_g B[z~_l,d,g] * S_l[d,g,k] M_l[k] = prod_{d=1..d_seed} phi_l[d,k] 4. Output. Embeddings are the sum across heads after a final linear projection to the transformer hidden width D: E_i = sum_{l=1..h} W_l M_l(z~_l), W_l in R^{r x D}. Parameter count: k*b*d_seed + h*(d_seed^2 + 2*d_seed + d_seed*kappa*r + r*D). With the paper defaults (h=8, r=64, kappa=16, d_seed=128, k=3, b=59) and V=200,376 this is ~1.21M + 512*D: at D=2048, 2.25M parameters vs V*D = 410.4M for a dense input table, keeping the untied output head at under 0.2% total parameter overhead at 1.2B scale. Spline weights use the authors' stable multiplicative parameterization (1 + wd_i), initialized wd_i ~ N(0, 0.1) so phi ~ 1 at init. """ def __init__(self, config: NeoLLMConfig): super().__init__() self.config = config # The Triton package is optional. Keep the original reference path # as the default whenever the package or its forward implementation # is unavailable; callers can also disable the optional route per # instance without changing the model/state-dict contract. self.use_leviathan_triton = bool(_LEV_KERNEL_AVAILABLE) self.dynamics_samples = ( int(getattr(config, "dynamics_sample_tokens", 64)) if bool(getattr(config, "dynamics_metrics_enabled", False)) else 0 ) self._last_dynamics_metrics = None vocab_size = config.vocab_size hidden_size = config.hidden_size self.d_seed = config.generator_d_seed self.num_modes = config.generator_num_modes # h heads self.num_knots = config.generator_num_knots # kappa self.spline_degree = config.generator_spline_degree self.k = config.generator_k self.krank = getattr(config, "generator_krank", 64) # r self.hidden_size = hidden_size # Base-b decomposition: b = ceil(V^(1/k)) covers the vocabulary with # minimal unused capacity (paper: k=3, b=59 covers 200,376 with 205,379). b = math.ceil(vocab_size ** (1.0 / self.k)) self.b = b # -- Stage 1: shared codebooks C_1..C_k in R^{b x d_seed} ------------ self.codebooks = nn.Parameter(torch.empty(self.k, b, self.d_seed)) # Fixed knot grid over [0, 1] (kappa points), not learned. # Keep it in the checkpoint: Hugging Face may materialize # non-persistent buffers from meta storage without initialization. self.register_buffer( "knot_grid", torch.linspace(0.0, 1.0, self.num_knots), persistent=True, ) # -- Stage 2: per-head seed projections W_seed,l in R^{d_seed x d_seed} # (no bias; a per-head LayerNorm is applied before the sigmoid) self.head_proj_weight = nn.Parameter( torch.empty(self.num_modes, self.d_seed, self.d_seed) ) self.head_norm_weight = nn.Parameter(torch.ones(self.num_modes, self.d_seed)) self.head_norm_bias = nn.Parameter(torch.zeros(self.num_modes, self.d_seed)) self.head_norm_eps = 1e-5 # -- Stage 3: separable spline coefficients S_l in R^{d_seed x kappa x r} # Effective coefficient used in the tensor product is (1 + delta). self.head_spline_delta = nn.Parameter( torch.empty(self.num_modes, self.d_seed, self.num_knots, self.krank) ) # -- Stage 4: per-head output projections W_l in R^{r x D} ------------ self.head_out_weight = nn.Parameter( torch.empty(self.num_modes, self.krank, self.hidden_size) ) # Shared continuous coordinate for Leviathan-JTok/JTok-M. The # original JTok paper stores a separate V×D table per layer; NeoLLM # replaces that table with a compact surface over one Leviathan seed. # These bridge parameters are present only when JTok is requested, so # the default Leviathan checkpoint/state-dict contract is unchanged. if bool(getattr(config, "use_jtok", False)): self.jtok_seed_proj_weight = nn.Parameter( torch.empty(self.d_seed, self.d_seed) ) self.jtok_seed_proj_bias = nn.Parameter(torch.zeros(self.d_seed)) self.jtok_seed_norm_weight = nn.Parameter(torch.ones(self.d_seed)) self.jtok_seed_norm_bias = nn.Parameter(torch.zeros(self.d_seed)) self.jtok_seed_norm_eps = 1e-5 else: self.register_parameter("jtok_seed_proj_weight", None) self.register_parameter("jtok_seed_proj_bias", None) self.register_parameter("jtok_seed_norm_weight", None) self.register_parameter("jtok_seed_norm_bias", None) self.jtok_seed_norm_eps = 1e-5 def _base_k_decompose(self, token_ids: torch.Tensor) -> torch.Tensor: """ Deterministic base-b decomposition: i -> (i_0, ..., i_{k-1}). Maps token indices directly to codebook coordinates via arithmetic: token x -> (x // b^{k-1}, ..., x % b). """ ids = token_ids.long().clone() coords = torch.empty( *token_ids.shape, self.k, dtype=torch.long, device=token_ids.device, ) for r in range(self.k - 1, -1, -1): coords[..., r] = ids % self.b ids = ids // self.b return coords def _codebook_seed(self, token_ids: torch.Tensor) -> torch.Tensor: """Return the differentiable compositional Leviathan seed ``z``.""" orig_shape = token_ids.shape N = token_ids.numel() coords = self._base_k_decompose(token_ids).reshape(N, self.k) z = torch.zeros( N, self.d_seed, device=token_ids.device, dtype=self.codebooks.dtype, ) for r in range(self.k): z = z + self.codebooks[r][coords[:, r]] return z.reshape(*orig_shape, self.d_seed) def _jtok_coordinate_from_seed(self, z: torch.Tensor) -> torch.Tensor: r"""Map a raw Leviathan seed to the shared bounded JTok coordinate. The bridge is a single learned ``Linear → LayerNorm → sigmoid`` path shared by all decoder layers. It preserves the smooth compositional geometry of Leviathan while giving every JTok surface the same coordinate system. """ if self.jtok_seed_proj_weight is None: raise RuntimeError( "JTok geometry requested but `use_jtok` was not enabled at " "model construction time." ) orig_shape = z.shape[:-1] flat = z.reshape(-1, self.d_seed) y = F.linear( flat.to(self.jtok_seed_proj_weight.dtype), self.jtok_seed_proj_weight, self.jtok_seed_proj_bias, ).float() mean = y.mean(dim=-1, keepdim=True) var = y.var(dim=-1, keepdim=True, unbiased=False) y = (y - mean) * torch.rsqrt(var + self.jtok_seed_norm_eps) y = y * self.jtok_seed_norm_weight.float() + self.jtok_seed_norm_bias.float() y = torch.sigmoid(0.5 * y).clamp(0.0, 1.0) return y.to(z.dtype).reshape(*orig_shape, self.d_seed) @staticmethod def _normalize_bspline_basis(B: torch.Tensor) -> torch.Tensor: """ Explicitly normalize a quadratic B-spline basis across knot points. The finite-grid quadratic kernel does not sum to exactly 1 near the boundaries of [0, 1]. Dividing by the post-evaluation sum across knots keeps every per-dimension basis vector on a partition-of-unity scale. """ denom = B.sum(dim=-1, keepdim=True).clamp_min(1e-12) return B / denom def _bspline_basis(self, x_flat: torch.Tensor) -> torch.Tensor: """ Quadratic B-spline basis with fixed scalar scale (num_knots - 1). Args: x_flat: [N, d_seed], values in [0, 1]. Returns: [N, d_seed, num_knots] float32. """ scale = float(self.num_knots - 1) x32 = x_flat.float() x_e = x32.unsqueeze(-1) grid = self.knot_grid.float().view(1, 1, -1) d = (x_e - grid).abs() * scale B = torch.where( d < 0.5, 0.75 - d**2, torch.where(d < 1.5, 0.5 * (1.5 - d) ** 2, torch.zeros_like(d)), ) # [N, d_seed, num_knots] float32 return self._normalize_bspline_basis(B) def _tensor_product( self, B: torch.Tensor, spline_coeff: torch.Tensor, ) -> torch.Tensor: """ Rank-r separable tensor product over the d_seed dimensions. phi[n, d, k] = sum_g B[n, d, g] * coeff[d, g, k] M[n, k] = prod_d phi[n, d, k] Computed in log space with sign-parity tracking (KHRONOS) so the product of d_seed factors neither underflows nor loses sign. """ phi = torch.einsum( "ndg,dgk->ndk", B, spline_coeff, ) # [N, d_seed, krank] log_mag = torch.log(phi.abs() + 1e-9).sum(dim=1) # [N, krank] num_neg = (phi < 0).to(torch.int32).sum(dim=1) # [N, krank] prod_sign = 1.0 - 2.0 * (num_neg % 2).float() # [N, krank] return prod_sign * torch.exp(log_mag) # [N, krank] def _record_dynamics(self, token_ids, embedding, checkpoints, valid_mask, mask_token_id): seed, xhat, modes = checkpoints diag_mask = (token_ids.new_empty((0,), dtype=torch.bool) if valid_mask is None else valid_mask.reshape(-1).contiguous()) if mask_token_id is not None: unmasked = token_ids.reshape(-1).ne(int(mask_token_id)) diag_mask = unmasked if diag_mask.numel() == 0 else diag_mask & unmasked packed = leviathan_dynamics( embedding.detach(), seed, xhat, self.head_norm_weight.detach(), self.head_norm_bias.detach(), modes, diag_mask, self.dynamics_samples) self._last_dynamics_metrics = packed.detach().clone() # Keep this method traceable. The compiler-safe LEV entry point registers # an opaque custom op tagged cudagraph_unsafe: Dynamo keeps the LEV call # in the compiled graph while Inductor excludes only the unsafe Triton # launch from CUDA-graph capture. A CUDA launch error is allowed to # propagate instead of falling back to reference work on the same stream. def forward( self, token_ids: torch.Tensor, *, meap_mask_embedding: Optional[torch.Tensor] = None, meap_mask_token_id: Optional[int] = None, return_jtok_geometry: bool = False, dynamics_valid_mask: Optional[torch.Tensor] = None, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: """ Generate embeddings from discrete token indices. E_i = sum_{l=1..h} W_l M_l(z~_l) with z~_l = sigmoid(1/2 LN(W_seed,l z)). Args: token_ids: (batch, seq_len) or (seq_len,). Returns: embeddings [*token_ids.shape, hidden_size]. When ``return_jtok_geometry=True``, returns ``(embeddings, z_tilde)`` where ``z_tilde`` is the shared differentiable Leviathan/JTok coordinate with shape ``[*token_ids.shape, d_seed]``. """ collect_dynamics = bool(self.dynamics_samples and self.training and torch.is_grad_enabled()) dynamics_kwargs = ( {"return_dynamics_checkpoints": True} if collect_dynamics else {} ) has_meap = ( meap_mask_embedding is not None or meap_mask_token_id is not None ) if has_meap: if meap_mask_embedding is None or meap_mask_token_id is None: raise ValueError( "meap_mask_embedding and meap_mask_token_id must be " "provided together" ) if ( meap_mask_embedding.ndim != 1 or meap_mask_embedding.numel() != self.hidden_size ): raise ValueError( "meap_mask_embedding must have shape " f"({self.hidden_size},)" ) leviathan_kernel_ready = ( self.use_leviathan_triton and _LEV_KERNEL_AVAILABLE and _leviathan_embedding_compiler_safe is not None and _leviathan_supports is not None and token_ids.is_cuda and _leviathan_config_supported(self.config) and all( parameter.dtype == torch.bfloat16 for parameter in ( self.codebooks, self.head_proj_weight, self.head_norm_weight, self.head_norm_bias, self.head_spline_delta, self.head_out_weight, ) ) ) require_leviathan_triton = bool( getattr(self.config, "_require_leviathan_triton", False) ) or collect_dynamics if require_leviathan_triton and not leviathan_kernel_ready: raise RuntimeError( "This run requires the Triton Leviathan kernel, but the " "installed CCE package or the current geometry is not " "supported. Refusing a reference fallback." ) if ( return_jtok_geometry and str(getattr(self.config, "jtok_kernel_backend", "torch")) .strip() .lower() == "triton" and not (leviathan_kernel_ready and _LEV_GEOMETRY_KERNEL_AVAILABLE) ): raise RuntimeError( "JTok Triton requires a supported CUDA Leviathan kernel and " "the CCE geometry adapter; refusing a silent reference fallback." ) # Ordinary Leviathan calls keep their existing embedding-only custom # op. JTok additionally needs the raw compositional seed. The CCE # adapter keeps the legacy kernel for the embedding and exposes only a # small differentiable base-b/codebook bridge for the geometry; this # avoids both a full reference-Leviathan detour and a silent cut in the # JTok -> codebook gradient. if ( leviathan_kernel_ready ): params = { "codebooks": self.codebooks, "head_proj_weight": self.head_proj_weight, "head_norm_weight": self.head_norm_weight, "head_norm_bias": self.head_norm_bias, "head_spline_delta": self.head_spline_delta, "head_out_weight": self.head_out_weight, } try: if return_jtok_geometry: if not _LEV_GEOMETRY_KERNEL_AVAILABLE: raise RuntimeError( "JTok geometry requested with the Leviathan kernel, " "but the installed CCE package lacks the geometry " "adapter." ) result = _leviathan_embedding_with_seed_compiler_safe( token_ids, params, self.config, self.knot_grid, mask_embedding=meap_mask_embedding, mask_token_id=meap_mask_token_id, **dynamics_kwargs, ) if collect_dynamics: embedding, seed, checkpoints = result self._record_dynamics(token_ids, embedding, checkpoints, dynamics_valid_mask, meap_mask_token_id) else: embedding, seed = result embedding = embedding.reshape(*token_ids.shape, self.hidden_size) return embedding, self._jtok_coordinate_from_seed(seed) result = _leviathan_embedding_compiler_safe( token_ids, params, self.config, self.knot_grid, mask_embedding=meap_mask_embedding, mask_token_id=meap_mask_token_id, **dynamics_kwargs, ) if collect_dynamics: embedding, checkpoints = result self._record_dynamics(token_ids, embedding, checkpoints, dynamics_valid_mask, meap_mask_token_id) return embedding return result except (TypeError, ValueError, AttributeError): # Only validation/configuration failures use the reference # implementation. A CUDA launch error must propagate: doing # reference work after a failed launch can poison the next # CUDA-graph/FlashAttention partition. if require_leviathan_triton: raise pass target_dtype = self.codebooks.dtype orig_shape = token_ids.shape N = token_ids.numel() # -- Stage 1: compositional codebook indexing ------------------------- z = self._codebook_seed(token_ids).reshape(N, self.d_seed) # -- Stages 2-4: per-head projection, spline expansion, sum ---------- e = torch.zeros( N, self.hidden_size, device=token_ids.device, dtype=target_dtype ) for m in range(self.num_modes): proj_w = self.head_proj_weight[m] # [d_seed, d_seed] zh = F.linear( z.to(dtype=proj_w.dtype, device=proj_w.device), proj_w, ) # [N, d_seed] zh = zh.float() # Per-head LayerNorm (independent weight/bias per head) norm_w = self.head_norm_weight[m].float() norm_b = self.head_norm_bias[m].float() mean = zh.mean(dim=-1, keepdim=True) var = zh.var(dim=-1, keepdim=True, unbiased=False) zh = (zh - mean) / (var + self.head_norm_eps).sqrt() zh = zh * norm_w + norm_b # Bounded latent coordinate: sigmoid(1/2 * LN(W_seed z)) in [0,1] zh = torch.sigmoid(zh / 2.0).clamp(0.0, 1.0) # [N, d_seed] # Separable quadratic B-spline basis + rank-r tensor product B = self._bspline_basis(zh) # [N, d_seed, num_knots] float32 modes = self._tensor_product( B, 1.0 + self.head_spline_delta[m].float(), ) # [N, krank] e = e + modes.to(self.head_out_weight.dtype) @ self.head_out_weight[m] e = e.reshape(*orig_shape, self.hidden_size) if has_meap: e = torch.where( token_ids.eq(int(meap_mask_token_id)).unsqueeze(-1), meap_mask_embedding.to(device=e.device, dtype=e.dtype), e, ) if return_jtok_geometry: z_jtok = self._jtok_coordinate_from_seed( z.reshape(*orig_shape, self.d_seed) ) return e, z_jtok return e # ==================== LEVIATHAN-JTOK / JTOK-M MODULATION ==================== class LeviathanJTok(nn.Module): r"""Pure-PyTorch Leviathan-coupled implementation of JTok and JTok-M. ``use_jtok=True, use_jtokm=False`` activates JTok. Enabling ``use_jtokm=True`` upgrades the same module to JTok-M; the configuration validates that JTok and the Leviathan generator are active first. The original JTok paper stores a token table per decoder layer. NeoLLM keeps the JTok modulation/routing equations but replaces every table row with a continuous CP-separable B-spline surface over the shared Leviathan-derived coordinate ``z_tilde``: .. math:: \hat{\Delta m}_x^l = \Delta m_x^l \odot (1+s^l\odot\operatorname{Norm}(S_l(\tilde z_x))) JTok-M evaluates the selected Top-K surfaces and injects their normalized mixture as ``Delta_m + Delta_r`` with the paper's ``1/sqrt(2L)`` depth factor. The default implementation is reference PyTorch. Setting the runtime-only ``config.jtok_kernel_backend`` to ``"triton"`` dispatches through the optional CCE JTok/JTok-M adapter while preserving this module's parameters, state-dict layout, and decoder return contract. The existing optional Leviathan embedding kernel remains a separate path for ordinary embedding calls. The Triton choice is strict: if the adapter is unavailable, forward raises instead of silently changing the training or benchmark path back to Torch. The temporary ``surface_chunk_tokens`` value bounds inference/evaluation work tiles only. It is not a context-length, batch-size, or VRAM limit. """ def __init__(self, config: NeoLLMConfig, layer_idx: int): super().__init__() self.layer_idx = int(layer_idx) self.use_mixture = bool(getattr(config, "use_jtokm", False)) self.dynamics_samples = ( int(getattr(config, "dynamics_sample_tokens", 64)) if bool(getattr(config, "dynamics_metrics_enabled", False)) else 0 ) self._last_dynamics_metrics = None self.jtok_kernel_backend = str( getattr(config, "jtok_kernel_backend", "torch") ).strip().lower() if self.jtok_kernel_backend not in {"torch", "triton"}: raise ValueError( "`jtok_kernel_backend` must be either `torch` or `triton`." ) if self.dynamics_samples and self.jtok_kernel_backend != "triton": raise RuntimeError("Native dynamics metrics require the strict Triton JTok backend") if self.jtok_kernel_backend == "triton" and not _JTOK_KERNEL_AVAILABLE: raise RuntimeError( "`jtok_kernel_backend='triton'` was requested, but the " "installed CCE package does not expose `apply_neollm_jtok`. " "Install the compatible JTok/JTok-M CCE implementation or " "set `jtok_kernel_backend='torch'`." ) self.hidden_size = int(config.hidden_size) self.d_seed = int(config.generator_d_seed) self.num_knots = int(config.generator_num_knots) self.num_modes = int(getattr(config, "jtok_num_modes", 4)) self.num_experts = ( int(getattr(config, "jtok_num_experts", 5)) if self.use_mixture else 1 ) self.top_k = ( int(getattr(config, "jtok_top_k", 2)) if self.use_mixture else 1 ) self.norm_eps = float(getattr(config, "jtok_norm_eps", 1e-6)) # Work-tile size only. The optional attribute is intentionally read # with getattr so older config.json files remain loadable. self.surface_chunk_tokens = max( int(getattr(config, "jtok_surface_chunk_tokens", 2048)), 1 ) self.jtokm_residual_scale = 1.0 / math.sqrt( 2.0 * float(config.num_hidden_layers) ) self.register_buffer( "knot_grid", torch.linspace(0.0, 1.0, self.num_knots), persistent=True, ) # Continuous replacement for the paper's per-token tables. These # remain ordinary Parameters so both their surface and the shared # Leviathan codebooks receive gradients during pretraining. self.spline_coeff = nn.Parameter( torch.empty( self.num_experts, self.num_modes, self.d_seed, self.num_knots, ) ) self.W_out = nn.Parameter( torch.empty(self.num_experts, self.num_modes, self.hidden_size) ) self.W_res = nn.Parameter( torch.empty(self.num_experts, self.d_seed, self.hidden_size) ) self.scaler = nn.Parameter(torch.ones(self.hidden_size)) self.router = ( nn.Linear(self.hidden_size, self.num_experts, bias=False) if self.use_mixture else None ) def _bspline_basis(self, z_flat: torch.Tensor) -> torch.Tensor: """Evaluate the normalized quadratic B-spline basis in FP32.""" scale = float(self.num_knots - 1) z32 = z_flat.float().unsqueeze(-1) grid = self.knot_grid.float().view(1, 1, -1) d = (z32 - grid).abs() * scale B = torch.where( d < 0.5, 0.75 - d.square(), torch.where( d < 1.5, 0.5 * (1.5 - d).square(), torch.zeros_like(d), ), ) return B / B.sum(dim=-1, keepdim=True).clamp_min(1e-12) def _eval_surfaces( self, z_flat: torch.Tensor, target_dtype: torch.dtype, ) -> torch.Tensor: """Evaluate all continuous surfaces, returning ``[N, E, D]``.""" B = self._bspline_basis(z_flat) phi = torch.einsum("emrg,nrg->nemr", self.spline_coeff.float(), B) log_mag = torch.log(phi.abs() + 1e-9).sum(dim=-1) num_neg = (phi < 0).to(torch.int32).sum(dim=-1) prod_sign = 1.0 - 2.0 * (num_neg % 2).float() modes = (prod_sign * torch.exp(log_mag)).to(target_dtype) out_modes = torch.einsum( "nem,emd->ned", modes, self.W_out.to(target_dtype) ) out_res = torch.einsum( "nr,erd->ned", z_flat.to(target_dtype), self.W_res.to(target_dtype), ) return out_modes + out_res def _eval_selected_surfaces( self, z_flat: torch.Tensor, expert_idx: torch.Tensor, target_dtype: torch.dtype, ) -> torch.Tensor: """Evaluate only the selected expert surfaces for each token.""" if expert_idx.ndim != 2: raise ValueError("expert_idx must have shape [tokens, selected_experts]") N, selected_k = expert_idx.shape basis_all = self._bspline_basis(z_flat) selected_flat = torch.zeros( N * selected_k, self.hidden_size, dtype=target_dtype, device=z_flat.device, ) # Top-K indices are unique per token, so each flattened destination is # written once. This avoids materializing [N,K,M,D] or [N,K,R,D]. for expert_idx_value in range(self.num_experts): token_idx, slot_idx = torch.where(expert_idx == expert_idx_value) flat_idx = token_idx * selected_k + slot_idx z_selected = z_flat[token_idx] B = basis_all[token_idx] coeff = self.spline_coeff[expert_idx_value].float() phi = torch.einsum("sdg,mdg->smd", B, coeff) log_mag = torch.log(phi.abs() + 1e-9).sum(dim=-1) num_neg = (phi < 0).to(torch.int32).sum(dim=-1) prod_sign = 1.0 - 2.0 * (num_neg % 2).float() modes = prod_sign * torch.exp(log_mag) values = modes.to(target_dtype) @ self.W_out[ expert_idx_value ].to(target_dtype) values = values + z_selected.to(target_dtype) @ self.W_res[ expert_idx_value ].to(target_dtype) selected_flat = selected_flat.index_copy(0, flat_idx, values) return selected_flat.reshape(N, selected_k, self.hidden_size) def _eval_selected_surfaces_tiled( self, z_flat: torch.Tensor, expert_idx: torch.Tensor, target_dtype: torch.dtype, ) -> torch.Tensor: """Route-first surface evaluation with bounded temporary token tiles.""" if z_flat.shape[0] != expert_idx.shape[0]: raise ValueError("z_flat and expert_idx must contain the same token count") chunks = [] for start in range(0, z_flat.shape[0], self.surface_chunk_tokens): stop = min(start + self.surface_chunk_tokens, z_flat.shape[0]) chunks.append( self._eval_selected_surfaces( z_flat[start:stop], expert_idx[start:stop], target_dtype ) ) if not chunks: return z_flat.new_empty((0, expert_idx.shape[1], self.hidden_size)) return torch.cat(chunks, dim=0) def _eval_selected_surfaces_dense_tiled( self, z_flat: torch.Tensor, expert_idx: torch.Tensor, target_dtype: torch.dtype, ) -> torch.Tensor: """Compile-friendly selected evaluation over bounded dense tiles.""" if z_flat.shape[0] != expert_idx.shape[0]: raise ValueError("z_flat and expert_idx must contain the same token count") chunks = [] for start in range(0, z_flat.shape[0], self.surface_chunk_tokens): stop = min(start + self.surface_chunk_tokens, z_flat.shape[0]) surfaces = self._eval_surfaces(z_flat[start:stop], target_dtype) gather_idx = expert_idx[start:stop].unsqueeze(-1).expand( stop - start, expert_idx.shape[1], self.hidden_size ) chunks.append(surfaces.gather(1, gather_idx)) if not chunks: return z_flat.new_empty((0, expert_idx.shape[1], self.hidden_size)) return torch.cat(chunks, dim=0) def _l2_normalize(self, x: torch.Tensor) -> torch.Tensor: return x / (x.norm(dim=-1, keepdim=True) + self.norm_eps) def _valid_mask( self, valid_mask: Optional[torch.Tensor], N: int, device: torch.device, ) -> torch.Tensor: if valid_mask is None: return torch.ones(N, device=device, dtype=torch.bool) return valid_mask.reshape(N).to(device=device, dtype=torch.bool) def _jtokm_aux_stats( self, logits: torch.Tensor, topk_idx: torch.Tensor, mask: torch.Tensor, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """Return differentiable probability and hard Top-K count reductions.""" aux_sigmoid = torch.sigmoid(logits.float()) # Normalize only the auxiliary probability branch. The forward # mixture remains selected-only Sigmoid+TopK, as in the source design. prob_all = aux_sigmoid / aux_sigmoid.sum( dim=-1, keepdim=True ).clamp_min(self.norm_eps) valid32 = mask.to(prob_all.dtype).unsqueeze(-1) p_sum = (prob_all * valid32).sum(dim=0) onehot = torch.zeros_like(prob_all).scatter(1, topk_idx, 1.0) f_sum = (onehot * valid32).sum(dim=0).detach() T = mask.sum().to(dtype=torch.float32).detach() return p_sum, f_sum, T def forward( self, delta_m: torch.Tensor, z_tilde: torch.Tensor, *, router_state: Optional[torch.Tensor] = None, valid_mask: Optional[torch.Tensor] = None, compute_aux: bool = False, ) -> Tuple[ torch.Tensor, Optional[Tuple[torch.Tensor, torch.Tensor, torch.Tensor]], ]: """Apply JTok/JTok-M and optionally return training-only router stats.""" collect_dynamics = bool(self.dynamics_samples and self.training and torch.is_grad_enabled()) if self.jtok_kernel_backend == "triton": # Dispatch only: this branch does not change compiler mode, # CUDA-graph policy, training state, optimizer behavior, or any # model-global state. The adapter returns the same # ``(updated_delta, compact_router_stats)`` contract as the Torch # implementation, so the surrounding pretraining graph remains # unchanged when the backend is switched. if _apply_neollm_jtok is None: raise RuntimeError( "JTok Triton adapter became unavailable after module " "construction; refusing an implicit Torch fallback." ) result = _apply_neollm_jtok( self, delta_m, z_tilde, router_state=router_state, valid_mask=valid_mask, compute_aux=bool(compute_aux and self.use_mixture), backend="triton", **({"return_dynamics_checkpoints": True} if collect_dynamics else {}), ) if collect_dynamics: output, aux, checkpoints = result modes, routes, weights = checkpoints n = delta_m.numel() // self.hidden_size mask = (delta_m.new_empty((0,), dtype=torch.bool) if valid_mask is None else valid_mask.reshape(-1).to(dtype=torch.bool).contiguous()) packed = jtok_dynamics( delta_m.detach().reshape(n, self.hidden_size).contiguous(), output.detach().reshape(n, self.hidden_size).contiguous(), z_tilde.detach().reshape(n, self.d_seed).contiguous(), modes, routes, weights, mask, self.num_experts, self.dynamics_samples) self._last_dynamics_metrics = packed.detach().clone() return output, aux return result orig_shape = delta_m.shape N = delta_m.numel() // self.hidden_size dm = delta_m.reshape(N, self.hidden_size) z = z_tilde.reshape(N, self.d_seed) mask = self._valid_mask(valid_mask, N, dm.device) mask_f = mask.to(dm.dtype).unsqueeze(-1) if not self.use_mixture: surfaces = self._eval_surfaces(z, dm.dtype) token_direction = self._l2_normalize(surfaces[:, 0, :]) gate = 1.0 + self.scaler.to(dm.dtype) * token_direction gate = torch.where(mask.unsqueeze(-1), gate, torch.ones_like(gate)) return (dm * gate).reshape(orig_shape), None if router_state is None: raise ValueError("JTok-M requires a pre-attention `router_state`.") h = router_state.reshape(N, self.hidden_size) h = h * torch.rsqrt(h.square().mean(dim=-1, keepdim=True) + self.norm_eps) logits = self.router(h) topk_vals, topk_idx = torch.topk(logits, self.top_k, dim=-1) selected_prob = torch.sigmoid(topk_vals) weights = selected_prob / selected_prob.sum( dim=-1, keepdim=True ).clamp_min(self.norm_eps) # Match the Torch reference semantics: keep the small dense path for # gradient-enabled training and use route-first tiles for large eager # evaluation. Both routes implement the same surface equations. if N <= self.surface_chunk_tokens or torch.is_grad_enabled(): surfaces = self._eval_surfaces(z, dm.dtype) gather_idx = topk_idx.unsqueeze(-1).expand( N, self.top_k, self.hidden_size ) selected = surfaces.gather(1, gather_idx) elif getattr(torch.compiler, "is_compiling", lambda: False)(): selected = self._eval_selected_surfaces_dense_tiled( z, topk_idx, dm.dtype ) else: selected = self._eval_selected_surfaces_tiled( z, topk_idx, dm.dtype ) mixed = (weights.unsqueeze(-1) * selected).sum(dim=1) delta_r = ( self.jtokm_residual_scale * self.scaler.to(dm.dtype) * self._l2_normalize(mixed) ) delta_r = delta_r * mask_f shared_update = dm + delta_r aux_stats = self._jtokm_aux_stats(logits, topk_idx, mask) if compute_aux else None return shared_update.reshape(orig_shape), aux_stats def compute_jtokm_aux_loss( aux_stats_list: Tuple[Tuple[torch.Tensor, torch.Tensor, torch.Tensor], ...], *, n_e: int, top_k: int, weight: float, ) -> torch.Tensor: r"""Compute the JTok-M load-balancing objective from compact reductions. For each layer, ``p_i=P_i/T`` and ``f_i=n_i/(T*K)``; the returned value is the layer-average of ``weight * n_e * sum_i(p_i*f_i)``. ``P_i`` remains attached to the router graph while hard Top-K counts are detached. """ if not aux_stats_list: raise ValueError("JTok-M auxiliary loss requires non-empty router stats.") layer_losses = [] for p_sum, f_sum, T in aux_stats_list: denom_t = T.to(device=p_sum.device, dtype=torch.float32).clamp_min(1.0) p_i = p_sum.float() / denom_t f_i = f_sum.to(device=p_sum.device, dtype=torch.float32) / ( denom_t * float(top_k) ) layer_losses.append(float(n_e) * (p_i * f_i).sum()) mean_loss = torch.stack(layer_losses).mean() lambda_t = torch.as_tensor(weight, device=mean_loss.device, dtype=torch.float32) return lambda_t * mean_loss # ==================== ORIGINAL COMPONENTS ==================== class FANLayer(nn.Module): """ Fourier Analysis Network (FAN) layer. FANLayer'(X) = [cos(WpX) || sin(WpX) || (Wp¯X + Bp¯)] """ def __init__(self, hidden_size: int, fan_ratio: float = 0.25): super().__init__() self.hidden_size = hidden_size self.fan_ratio = fan_ratio output_dim = hidden_size + int(hidden_size * fan_ratio) self.p_output_dim = int(output_dim * fan_ratio) self.g_output_dim = output_dim - self.p_output_dim * 2 self.input_linear = nn.Linear( hidden_size, self.p_output_dim + self.g_output_dim, bias=True ) self._init_weights() def _init_weights(self): nn.init.normal_(self.input_linear.weight, mean=0.0, std=0.02) if self.input_linear.bias is not None: nn.init.zeros_(self.input_linear.bias) def forward( self, x: torch.Tensor, ) -> torch.Tensor: pg = self.input_linear(x) p, g = torch.split(pg, [self.p_output_dim, self.g_output_dim], dim=-1) cos_p = torch.cos(p) sin_p = torch.sin(p) return torch.cat([cos_p, sin_p, g], dim=-1) class StackMemory(nn.Module): """ Differentiable hidden-state stack used by STACKTRANS. Hybrid StackTrans geometry: use the released OLMo v3 low-rank projection geometry (global D -> H*d_s and global H*d_s -> D) while preserving the Llama v9_1 token-local stack state across decoder layers instead of collapsing to the final position. Shapes: hidden_states: [B, S, D] stack: [B, S, H, K, d_s] mask: [B, S, H, K] where ``H = num_mem_heads``, ``K = stack_slots``, and ``H*d_s = stack_d_model``. With the current configuration, ``24 = 4 heads x 6 stack dimensions``. """ def __init__(self, config: "NeoLLMConfig"): super().__init__() self.config = config self.num_mem_heads = config.num_mem_heads self.stack_slots = config.stack_slots if config.stack_d_model % self.num_mem_heads != 0: raise ValueError( "StackMemory requires stack_d_model to be divisible by " "num_mem_heads." ) # OLMo stack_v3 geometry: one global low-rank projection D -> H*d_s, # followed by a reshape into H memory heads. self.stack_dim = config.stack_d_model // self.num_mem_heads self.down_proj = nn.Linear(config.hidden_size, config.stack_d_model) self.up_proj = nn.Linear(config.stack_d_model, config.hidden_size) self.action_head = nn.Linear(config.stack_d_model, 3 * self.num_mem_heads) # A scalar bias is exactly softmax-invariant because the same value is # added to every stack slot. Keeping it would create a parameter tensor # with an analytically zero gradient, which is especially undesirable for # per-tensor GradientStabilizer scaling. self.gate_proj = nn.Linear(self.stack_dim, 1, bias=False) self.res_weight = nn.Parameter(torch.ones(1)) self.cache_size = getattr(config, "stack_memory_cache_size", 2048) self.cache_position = 0 self.enable_cache = False def reset_cache(self): self.cache_position = 0 def _vectorized_update(self, stack, mask, actions, k_values): # ``stack`` is already token-local: [B, S, H, K, r]. Unlike the old # NeoLLM path, there is no sequence broadcast from one shared stack. push_stack = torch.cat([k_values.unsqueeze(3), stack[:, :, :, :-1]], dim=3) push_mask = torch.cat( [torch.ones_like(mask[:, :, :, :1]), mask[:, :, :, :-1]], dim=3 ) pop_stack = torch.cat( [stack[:, :, :, 1:], torch.zeros_like(stack[:, :, :, :1])], dim=3 ) pop_mask = torch.cat( [mask[:, :, :, 1:], torch.zeros_like(mask[:, :, :, :1])], dim=3 ) # Keep stack values in the model dtype (normally BF16), but keep the # differentiable control mask in FP32. The latter is important because # the global-read mask bias can be extremely large (e.g. -1e9). action_weights_stack = actions.to(stack.dtype).unsqueeze(-1).unsqueeze(-1) action_weights_mask = actions.to(mask.dtype).unsqueeze(-1) stacks = torch.stack([push_stack, pop_stack, stack], dim=3) masks = torch.stack([push_mask, pop_mask, mask], dim=3) new_stack = (stacks * action_weights_stack).sum(dim=3) new_mask = (masks * action_weights_mask).sum(dim=3) return new_stack, new_mask def forward(self, hidden_states, stack, mask, token_mask=None): batch_size, seq_len, _ = hidden_states.shape # The stack occupancy is a differentiable control state. Keep it in # FP32 across layers/callers even when model activations are BF16. mask = mask.float() # OLMo stack_v3 value/action path: project the complete hidden state into # the total low-rank stack width, then split that representation by head. new_hidden_states = self.down_proj(hidden_states) action_logits = self.action_head(new_hidden_states) / math.sqrt(self.stack_dim) # Stack actions are control probabilities. Evaluate softmax in FP32 so # the recurrent mask is never quantized to BF16 before later layers use # it. Gradients still flow normally into the BF16 action_head weights. actions = F.softmax( action_logits.float().reshape( batch_size, seq_len, self.num_mem_heads, 3 ), dim=-1, ) k_values = new_hidden_states.reshape( batch_size, seq_len, self.num_mem_heads, self.stack_dim ) new_stack, new_mask = self._vectorized_update(stack, mask, actions, k_values) # The global-read control path is evaluated explicitly in FP32. This # preserves the original differentiable graph (no detach) while avoiding # BF16 quantization of both the soft mask and the very large mask bias. stack_fp32 = new_stack.float() gate_scores = F.linear( stack_fp32, self.gate_proj.weight.float(), bias=None, ).squeeze(-1) gate_weights = F.softmax( gate_scores + (1.0 - new_mask) * -320, dim=-1, ) memory_output = (stack_fp32 * gate_weights.unsqueeze(-1)).sum(dim=3) memory_output = memory_output.to(new_hidden_states.dtype) memory_output = memory_output.reshape(batch_size, seq_len, -1) memory_output = self.up_proj(memory_output) output = memory_output * self.res_weight + hidden_states # Compact StackTrans statistics. Entropy is evaluated in FP32 for # numerical stability; only its scalar mean keeps autograd because it is # used by the auxiliary objective. All other diagnostics are detached # before they leave this layer so monitoring does not retain extra graphs. actions_fp32 = actions.float() action_entropy_map = -( actions_fp32 * torch.log(actions_fp32.clamp_min(1e-12)) ).sum(dim=-1) # [B,S,H] action_max_map = actions_fp32.max(dim=-1).values gate_fp32 = gate_weights.float() gate_entropy_map = -( gate_fp32 * torch.log(gate_fp32.clamp_min(1e-12)) ).sum(dim=-1) # [B,S,H] gate_max_map = gate_fp32.max(dim=-1).values gate_effective_slots_map = gate_entropy_map.exp() mask_fp32 = new_mask.float() expected_depth_map = mask_fp32.sum(dim=-1) # [B,S,H] mask_fill_map = expected_depth_map / float(self.stack_slots) top_occupancy_map = mask_fp32[..., 0] if token_mask is None: action_entropy_mean = action_entropy_map.mean() action_probs_mean = actions_fp32.mean(dim=(0, 1, 2)) action_max_mean = action_max_map.mean() expected_depth_mean = expected_depth_map.mean() mask_fill_mean = mask_fill_map.mean() top_occupancy_mean = top_occupancy_map.mean() gate_entropy_mean = gate_entropy_map.mean() gate_max_mean = gate_max_map.mean() gate_effective_slots_mean = gate_effective_slots_map.mean() else: valid = token_mask.to( device=hidden_states.device, dtype=torch.float32 ).unsqueeze(-1) # [B,S,1] denom = (valid.sum() * float(self.num_mem_heads)).clamp_min(1.0) action_entropy_mean = (action_entropy_map * valid).sum() / denom action_probs_mean = ( actions_fp32 * valid.unsqueeze(-1) ).sum(dim=(0, 1, 2)) / denom action_max_mean = (action_max_map * valid).sum() / denom expected_depth_mean = (expected_depth_map * valid).sum() / denom mask_fill_mean = (mask_fill_map * valid).sum() / denom top_occupancy_mean = (top_occupancy_map * valid).sum() / denom gate_entropy_mean = (gate_entropy_map * valid).sum() / denom gate_max_mean = (gate_max_map * valid).sum() / denom gate_effective_slots_mean = ( gate_effective_slots_map * valid ).sum() / denom stack_metrics = torch.stack( ( action_entropy_mean, action_probs_mean[0].detach(), action_probs_mean[1].detach(), action_probs_mean[2].detach(), action_max_mean.detach(), expected_depth_mean.detach(), mask_fill_mean.detach(), top_occupancy_mean.detach(), gate_entropy_mean.detach(), gate_max_mean.detach(), gate_effective_slots_mean.detach(), ) ) if self.training and self.enable_cache: self._update_cache(k_values.detach(), actions.detach()) # Preserve the complete token-local stack for the next decoder layer. return output, new_stack, new_mask, stack_metrics def _update_cache(self, k_values, actions): seq_len = k_values.shape[1] if self.cache_position + seq_len <= self.cache_size: self.k_cache[self.cache_position : self.cache_position + seq_len] = ( k_values[0] ) self.action_cache[self.cache_position : self.cache_position + seq_len] = ( actions[0] ) self.cache_position += seq_len else: self.reset_cache() def step(self, hidden_state, stack, mask): """Single-token wrapper for the token-local stack contract.""" # ``forward`` expects an explicit sequence dimension. Generation callers # conventionally hold only the current token's [B,H,K,r] stack state. if stack.ndim == 4: stack = stack.unsqueeze(1) if mask.ndim == 3: mask = mask.unsqueeze(1) output, new_stack, new_mask, _stack_metrics = self.forward( hidden_state.unsqueeze(1), stack, mask ) return output.squeeze(1), new_stack.squeeze(1), new_mask.squeeze(1) class LNS(nn.Module): """ LayerNorm Scaling: applies 1/√ℓ to suppress variance growth with depth. From "The Curse of Depth in Large Language Models". """ def __init__(self, layer_idx: int): super().__init__() self.layer_idx = max(layer_idx + 1, 1) # Keep the depth factor as a non-persistent buffer so Dynamo sees it # as tensor data shared by the compiled module contract, rather than a # different Python float guard for every decoder layer. It is not a # learned weight and is intentionally excluded from checkpoints. self.register_buffer( "scale", torch.tensor(1.0 / math.sqrt(self.layer_idx), dtype=torch.float32), persistent=False, ) def forward(self, x: torch.Tensor) -> torch.Tensor: # Match the previous scalar-multiply dtype while keeping the value on # the activation device. The buffer is cast with the model and the # ``to`` is a no-op on the normal steady-state path. return torch.ops.aten.mul.Tensor( x, self.scale.to(device=x.device, dtype=x.dtype), ) class GPAS(nn.Module): """Gradient-Preserving Activation Scaling.""" def __init__(self, d_model: int): super().__init__() self.d_model = d_model self.alpha = nn.Parameter(torch.zeros(1)) def forward( self, x: torch.Tensor, ) -> torch.Tensor: silu_alpha = F.silu(self.alpha) subtracted = silu_alpha * x.detach() return x - subtracted def _make_norm(dim: int, eps: float) -> nn.Module: """Build the active Transformer normalization module. The dynamic normalization path has been removed, so every backbone, Q/K, final, and MEA normalization site uses standard RMSNorm. """ return nn.RMSNorm(dim, eps=eps) def _apply_norm( norm: nn.Module, x: torch.Tensor, ) -> torch.Tensor: """Apply RMSNorm and optionally record the normalized output.""" output = norm(x) return output # ==================== ROTARY EMBEDDING ==================== class NeoLLMRotaryEmbedding(nn.Module): inv_freq: torch.Tensor def __init__(self, config: NeoLLMConfig, device=None): super().__init__() self.max_seq_len_cached = config.max_position_embeddings self.original_max_seq_len = config.max_position_embeddings self.config = config self.rope_type = "default" if ( hasattr(config, "rope_scaling") and config.rope_scaling is not None and isinstance(config.rope_scaling, dict) ): rope_type = config.rope_scaling.get( "rope_type", config.rope_scaling.get("type") ) if rope_type and rope_type in ROPE_INIT_FUNCTIONS: self.rope_type = rope_type rope_init_fn = self.compute_default_rope_parameters if self.rope_type != "default": rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type] inv_freq, self.attention_scaling = rope_init_fn(self.config, device) self.register_buffer("inv_freq", inv_freq, persistent=False) self.register_buffer("original_inv_freq", inv_freq.clone(), persistent=False) def _apply(self, fn): # Positional frequencies are an FP32 numerical contract for both the # PyTorch and fused REPO-GRAPE paths. Preserve that contract across # model.to(bfloat16)/model.to(float16) instead of allocating a cast in # every decoder layer. The buffers remain non-persistent as before. super()._apply(fn) for name in ("inv_freq", "original_inv_freq"): value = self._buffers.get(name) if value is not None: self._buffers[name] = value.float() return self @staticmethod def compute_default_rope_parameters( config: NeoLLMConfig = None, device: Optional["torch.device"] = None, seq_len: int = None, ) -> tuple["torch.Tensor", float]: base = config.rope_theta dim = ( getattr(config, "head_dim", None) or config.hidden_size // config.num_attention_heads ) dim = int(dim * getattr(config, "partial_rotary_factor", 1.0)) inv_freq = 1.0 / ( base ** ( torch.arange(0, dim, 2, dtype=torch.int64).to( device=device, dtype=torch.float ) / dim ) ) return inv_freq, 1.0 @torch.no_grad() @dynamic_rope_update def forward(self, x, position_ids): if position_ids.dim() == 1: position_ids = position_ids.unsqueeze(0) B = x.shape[0] if position_ids.shape[0] != B: position_ids = position_ids.expand(B, -1) device_type = ( x.device.type if isinstance(x.device.type, str) and x.device.type != "mps" else "cpu" ) if self.inv_freq.device.type == "meta": inv_freq_data, _ = self.compute_default_rope_parameters( self.config, device=x.device ) self.register_buffer("inv_freq", inv_freq_data, persistent=False) self.register_buffer( "original_inv_freq", inv_freq_data.clone(), persistent=False ) inv_freq = self.inv_freq.to(device=x.device, dtype=torch.float32) with torch.autocast(device_type=device_type, enabled=False): freqs = position_ids.to(dtype=torch.float32).unsqueeze( -1 ) * inv_freq.unsqueeze(0).unsqueeze(0) emb = torch.cat((freqs, freqs), dim=-1) cos = emb.cos() * self.attention_scaling sin = emb.sin() * self.attention_scaling return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype) return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype) def rotate_half(x): x1 = x[..., : x.shape[-1] // 2] x2 = x[..., x.shape[-1] // 2 :] return torch.cat((-x2, x1), dim=-1) def linear_clipping(x: torch.Tensor) -> torch.Tensor: """ Piecewise-linear activation for Affine-Scaled Attention scaling factors. Avoids the saturation problem of sigmoid, which collapses most outputs toward 0 or 1 and loses the intermediate scaling range the model needs for fine-grained per-query attention modulation. f(x) = 0 if x ≤ -5 = 0.1·x + 0.5 if -5 < x < 5 = 1 if x ≥ 5 Equivalent to: clamp(0.1·x + 0.5, 0, 1). Output range: [0, 1]. Gradient: 0.1 across the entire non-saturated region. Reference: Bae et al. (2026), Affine-Scaled Attention §6. """ return torch.clamp(0.1 * x + 0.5, 0.0, 1.0) def apply_rotary_pos_emb(q, k, cos, sin, position_ids=None, unsqueeze_dim=1): cos = cos.unsqueeze(unsqueeze_dim) sin = sin.unsqueeze(unsqueeze_dim) rotary_dim = cos.shape[-1] q_rot, q_pass = q[..., :rotary_dim], q[..., rotary_dim:] k_rot, k_pass = k[..., :rotary_dim], k[..., rotary_dim:] q_embed = (q_rot * cos) + (rotate_half(q_rot) * sin) k_embed = (k_rot * cos) + (rotate_half(k_rot) * sin) return torch.cat([q_embed, q_pass], dim=-1), torch.cat([k_embed, k_pass], dim=-1) def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor: batch, num_key_value_heads, slen, head_dim = hidden_states.shape if n_rep == 1: return hidden_states hidden_states = hidden_states[:, :, None, :, :].expand( batch, num_key_value_heads, n_rep, slen, head_dim ) return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim) def _flash_target_dtype( query: torch.Tensor, module: nn.Module, ) -> Optional[torch.dtype]: """Return the half dtype required by FlashAttention for FP32 Q/K states. PyTorch RMSNorm intentionally returns FP32 under BF16 autocast. IHA's pseudo-head einsums cast those states back through autocast, while the standard path has no equivalent operation. Resolve the target here so both routes enter every FlashAttention backend with the same dtype instead of relying on a backend-specific compatibility cast. """ if query.dtype != torch.float32: return None device_type = query.device.type if torch.is_autocast_enabled(device_type): return torch.get_autocast_dtype(device_type) if hasattr(module.config, "_is_quantized"): return module.config.dtype return next( layer.weight.dtype for layer in module.modules() if isinstance(layer, nn.Linear) ) def neollm_flash_attention_forward( module: nn.Module, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, attention_mask: Optional[torch.Tensor], dropout: float = 0.0, scaling: Optional[float] = None, sliding_window: Optional[int] = None, softcap: Optional[float] = None, is_causal: Optional[bool] = None, s_aux: Optional[torch.Tensor] = None, **kwargs: Unpack[TransformersKwargs], ): """Compile-stable HF FlashAttention bridge for NeoLLM. This mirrors the Transformers FA2/FA3 integration but deliberately does not forward ``module.layer_idx``. That integer is not consumed by the current FA2/FA3 implementation (it is discarded through ``**kwargs``), yet Dynamo guards its value and otherwise compiles one copy per decoder layer. Keeping the bridge local preserves the public integer attribute for model tooling, checkpoints, and diagnostics without making it part of the attention graph signature. """ if kwargs.get("output_attentions", False): logger.warning_once( "Flash Attention does not support `output_attentions=True`. " "Please set attention to `eager` to request attention weights." ) seq_len = query.shape[2] if any(dim == 0 for dim in query.shape): raise ValueError( "FlashAttention does not support a query tensor with a zero dimension." ) query = query.transpose(1, 2) key = key.transpose(1, 2) value = value.transpose(1, 2) is_causal = is_causal if is_causal is not None else module.is_causal attn_output = _flash_attention_forward( query, key, value, attention_mask, query_length=seq_len, is_causal=is_causal, dropout=dropout, softmax_scale=scaling, sliding_window=sliding_window, softcap=softcap, use_top_left_mask=_NEOLLM_FA_USE_TOP_LEFT_MASK, target_dtype=None, # NeoLLMAttention normalizes Q/K/V before dispatch. attn_implementation=module.config._attn_implementation, s_aux=s_aux.to(query.dtype) if s_aux is not None else None, **kwargs, ) return attn_output, None def causal_first_difference(x: torch.Tensor) -> torch.Tensor: return x - F.pad(x[..., :-1, :], (0, 0, 1, 0)) def rms_key_unit_norm(x: torch.Tensor, eps: float) -> torch.Tensor: return F.normalize(x.float(), p=2, dim=-1, eps=eps) * math.sqrt(x.shape[-1]) def infer_key_validity( attention_mask: Optional[torch.Tensor], seq_len: int, num_heads: int ) -> Optional[torch.Tensor]: if attention_mask is None: return None if attention_mask.ndim == 2: if attention_mask.shape[-1] != seq_len: return None valid = attention_mask.to(dtype=torch.bool).unsqueeze(1) # [B, 1, S] else: if attention_mask.ndim != 4: return None if attention_mask.shape[-2] != seq_len or attention_mask.shape[-1] != seq_len: return None diag = attention_mask.diagonal(dim1=-2, dim2=-1) if attention_mask.dtype == torch.bool: valid = diag else: valid = torch.isfinite(diag) & (diag == 0) if valid.shape[1] == 1 and num_heads != 1: valid = valid.expand(-1, num_heads, -1) elif valid.shape[1] != num_heads: valid = valid[:, :1, :].expand(-1, num_heads, -1) return valid def _expand_pair_valid_mask( attention_mask: Optional[torch.Tensor], batch_size: int, num_heads: int, query_len: int, key_len: int, ) -> Optional[torch.Tensor]: """ Return a boolean [B,H,Sq,Sk] mask for positions that attention may read. Supported contracts used in this file: • 2D padding mask [B,Sk], with truthy values for valid keys. • 4D additive mask [B,1|H,Sq,Sk], with 0 on valid pairs and -inf/large negative values on invalid causal, local or padding pairs. """ if attention_mask is None: return None if attention_mask.ndim == 2: if attention_mask.shape[-1] < key_len: return None valid = attention_mask[:, :key_len].to(dtype=torch.bool) # [B,Sk] valid = valid[:, None, None, :].expand( batch_size, num_heads, query_len, key_len ) return valid if attention_mask.ndim != 4: return None if attention_mask.shape[-2] < query_len or attention_mask.shape[-1] < key_len: return None mask = attention_mask[:, :, :query_len, :key_len] valid = mask if mask.dtype == torch.bool else mask == 0 if valid.shape[1] == 1 and num_heads != 1: valid = valid.expand(-1, num_heads, -1, -1) elif valid.shape[1] != num_heads: valid = valid[:, :1, :, :].expand(-1, num_heads, -1, -1) return valid def _masked_value_sum_from_pairs( value_expanded: torch.Tensor, pair_valid: torch.Tensor, ) -> torch.Tensor: """Compute Σ_{j∈A(i)} V_j from an explicit pair-valid mask.""" return torch.matmul(pair_valid.to(value_expanded.dtype), value_expanded) def _causal_value_sum( value_states: torch.Tensor, query_len: int, window_left: Optional[int] = None, ) -> torch.Tensor: """ Compute Σ_{j∈A(i)} V_j for causal or sliding-window causal attention. value_states: [B,H,Sk,d]. The returned tensor is [B,H,Sq,d]. For cached decoding, Sq may be smaller than Sk; query rows are aligned to the last Sq key positions, matching causal self-attention with cache. The sliding-window path uses prefix differences with slices, not padded gathers. This keeps the operation exact while avoiding extra index tensors and repeated index_select calls in the hot path. """ key_len = value_states.shape[2] prefix = value_states.cumsum(dim=2) # [B,H,Sk,d] if window_left is None or window_left < 0 or int(window_left) >= key_len: sums_all = prefix else: left = int(window_left) sums_all = prefix.clone() if left + 1 < key_len: # For i > left: # sum_{j=i-left}^{i} V_j = prefix[i] - prefix[i-left-1]. # The first left+1 rows are already the causal prefixes. sums_all[:, :, left + 1 :, :] = ( prefix[:, :, left + 1 :, :] - prefix[:, :, : -(left + 1), :] ) if query_len == key_len: return sums_all start = max(key_len - query_len, 0) return sums_all[:, :, start : start + query_len, :] def _segmented_causal_value_sum( value_states: torch.Tensor, segment_lengths: torch.Tensor, window_left: Optional[int] = None, ) -> torch.Tensor: """ Compute segmented causal/window value sums without a Python loop over documents. Used only for packed IHA where B=1 and position_ids reset. value_states: [1,H,Sk,d] segment_lengths: [num_segments] lengths in the same expanded sequence axis return: [1,H,Sk,d] """ key_len = value_states.shape[2] if key_len == 0: return value_states lengths = segment_lengths.to(device=value_states.device, dtype=torch.long) seg_ends = lengths.cumsum(0) seg_starts = seg_ends - lengths seg_start_per_pos = torch.repeat_interleave(seg_starts, lengths) pos = torch.arange(key_len, device=value_states.device, dtype=torch.long) if window_left is None or window_left < 0: start_idx = seg_start_per_pos else: start_idx = torch.maximum(seg_start_per_pos, pos - int(window_left)) end_idx = pos + 1 prefix = F.pad(value_states.cumsum(dim=2), (0, 0, 1, 0)) # [1,H,Sk+1,d] end_vals = prefix.index_select(2, end_idx) start_vals = prefix.index_select(2, start_idx) return end_vals - start_vals def _affine_valid_value_sum( value_expanded: torch.Tensor, attention_mask: Optional[torch.Tensor], query_len: int, sliding_window: Optional[Union[int, Tuple[int, int]]] = None, ) -> torch.Tensor: """ Compute the β term support S_i = Σ_{j∈A(i)} V_j for Affine attention. It first respects an explicit 4D additive mask exactly. Otherwise it uses causal prefix sums, optionally restricted by a left sliding-window size. """ batch_size, num_heads, key_len, _ = value_expanded.shape pair_valid = None if attention_mask is not None and attention_mask.ndim == 4: pair_valid = _expand_pair_valid_mask( attention_mask, batch_size, num_heads, query_len, key_len ) if pair_valid is not None: return _masked_value_sum_from_pairs(value_expanded, pair_valid) # 2D padding masks zero invalid keys before the prefix/window sum. if attention_mask is not None and attention_mask.ndim == 2: if attention_mask.shape[-1] >= key_len: valid = attention_mask[:, :key_len].to(value_expanded.dtype) value_expanded = value_expanded * valid[:, None, :, None] window_left = None if isinstance(sliding_window, tuple): window_left = int(sliding_window[0]) if len(sliding_window) > 0 else None elif sliding_window is not None: window_left = int(sliding_window) return _causal_value_sum(value_expanded, query_len, window_left) def head_linear_compose( hidden_states: torch.Tensor, mixing_matrix: torch.Tensor ) -> torch.Tensor: return torch.einsum( "bhtd,hk->bktd", hidden_states, mixing_matrix.to(device=hidden_states.device, dtype=hidden_states.dtype), ) def head_linear_compose_pseudo( hidden_states: torch.Tensor, mixing_matrix: torch.Tensor, num_pseudo_heads: int, ) -> torch.Tensor: """ Applies the same MEA head mixing independently inside each IHA pseudo-slot. Layout convention: hidden_states: [B, H_in * P, S, d] with flattened order (head, pseudo) mixing_matrix: [H_in, H_out] returns: [B, H_out * P, S, d] """ if num_pseudo_heads <= 1: return head_linear_compose(hidden_states, mixing_matrix) batch, total_heads, slen, head_dim = hidden_states.shape if total_heads % num_pseudo_heads != 0: raise ValueError( f"IHA+MEA expected total_heads divisible by P, got " f"{total_heads} vs P={num_pseudo_heads}." ) num_component_heads = total_heads // num_pseudo_heads if mixing_matrix.shape[0] != num_component_heads: raise ValueError( f"IHA+MEA expected {num_component_heads} component heads for MEA, " f"got mixing_matrix.shape[0]={mixing_matrix.shape[0]}." ) hidden_states = hidden_states.reshape( batch, num_component_heads, num_pseudo_heads, slen, head_dim ) mixed = torch.einsum( "bhpsd,hk->bkpsd", hidden_states, mixing_matrix.to(device=hidden_states.device, dtype=hidden_states.dtype), ) return mixed.reshape( batch, mixing_matrix.shape[1] * num_pseudo_heads, slen, head_dim ) class MEAHeadRMSNorm(nn.Module): """MEA head-level RMS normalization grouped by KV structure (GQA-aware).""" def __init__( self, num_heads: int, head_dim: int, num_kv_groups: int, eps: float = 1e-6, ): super().__init__() self.num_heads = num_heads self.head_dim = head_dim self.num_kv_groups = num_kv_groups self.num_kv_heads = num_heads // num_kv_groups self.group_dim = num_kv_groups * head_dim self.norm = _make_norm(self.group_dim, eps=eps) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: batch, seq_len, num_heads, head_dim = hidden_states.shape if num_heads != self.num_heads or head_dim != self.head_dim: raise ValueError( f"MEAHeadRMSNorm expected ({self.num_heads}, {self.head_dim}), " f"received ({num_heads}, {head_dim})" ) grouped = hidden_states.reshape( batch, seq_len, self.num_kv_heads, self.group_dim ) return self.norm(grouped).reshape(batch, seq_len, num_heads, head_dim) def eager_attention_forward( module: nn.Module, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, attention_mask: Optional[torch.Tensor], scaling: float, dropout: float = 0.0, **kwargs: Unpack[TransformersKwargs], ): key_states = repeat_kv(key, module.num_key_value_groups) value_states = repeat_kv(value, module.num_key_value_groups) attn_weights = torch.matmul(query, key_states.transpose(2, 3)) * scaling if attention_mask is not None: attn_weights = attn_weights + attention_mask[:, :, :, : key_states.shape[-2]] attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to( query.dtype ) attn_output = torch.matmul(attn_weights, value_states).transpose(1, 2).contiguous() attn_output = nn.functional.dropout( attn_output, p=dropout, training=module.training ) return attn_output, attn_weights def affine_scaled_eager_attention_forward( module: nn.Module, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, attention_mask: Optional[torch.Tensor], scaling: float, alpha: torch.Tensor, beta: torch.Tensor, dropout: float = 0.0, **kwargs: Unpack[TransformersKwargs], ): """ Affine-Scaled Attention (eager path). Replaces the standard weighted sum softmax(QK^T/√dk) V with: [α(X) · softmax(QK^T/√dk) + β(X)] V α is a per-head, per-query input-dependent scale in [0, 1]. β is an input-dependent bias that compensates for deviations of α from its running average, preventing the effective attention mass from collapsing. Both α and β are computed in NeoLLMAttention.forward and passed in; this function only performs the affine reweighting and the value aggregation. The existing Gated Attention gate (applied post-SDPA to the concatenated output before o_proj) is orthogonal to this and is not modified here. Reference: Bae et al. (2026), Affine-Scaled Attention, Eq. 6–8. Args: alpha: [batch, num_heads, seq_q, 1] — input-dependent scale per query beta: [batch, num_heads, seq_q, 1] — input-dependent bias per query """ key_states = repeat_kv(key, module.num_key_value_groups) value_states = repeat_kv(value, module.num_key_value_groups) attn_weights = torch.matmul(query, key_states.transpose(2, 3)) * scaling pair_valid = _expand_pair_valid_mask( attention_mask, batch_size=query.shape[0], num_heads=key_states.shape[1], query_len=query.shape[-2], key_len=key_states.shape[-2], ) if attention_mask is not None: if attention_mask.ndim == 4: attn_weights = ( attn_weights + attention_mask[:, :, : query.shape[-2], : key_states.shape[-2]] ) elif attention_mask.ndim == 2: key_valid = attention_mask[:, : key_states.shape[-2]].to(dtype=torch.bool) additive = torch.zeros_like(attn_weights) additive = additive.masked_fill( ~key_valid[:, None, None, :], torch.finfo(attn_weights.dtype).min ) attn_weights = attn_weights + additive attn_weights_softmax = nn.functional.softmax( attn_weights, dim=-1, dtype=torch.float32 ).to(query.dtype) if pair_valid is not None: attn_weights_softmax = attn_weights_softmax.masked_fill(~pair_valid, 0.0) # Affine reweighting over valid keys only. β must obey the same causal, # local and padding mask as softmax; otherwise invalid V_j would leak in. # Shapes: α, β are [B, H, S_q, 1], weights are [B, H, S_q, S_k]. attn_weights_affine = alpha * attn_weights_softmax + beta if pair_valid is not None: attn_weights_affine = attn_weights_affine.masked_fill(~pair_valid, 0.0) attn_output = ( torch.matmul(attn_weights_affine, value_states).transpose(1, 2).contiguous() ) attn_output = nn.functional.dropout( attn_output, p=dropout, training=module.training ) return attn_output, attn_weights_affine def affine_scaled_flash_attention_forward( module: nn.Module, query: torch.Tensor, key: torch.Tensor, value: torch.Tensor, attention_mask: Optional[torch.Tensor], scaling: float, alpha: torch.Tensor, beta: torch.Tensor, dropout: float = 0.0, **kwargs: Unpack[TransformersKwargs], ): """ Affine-Scaled Attention — flash/sdpa path. Mathematical decomposition of [α·softmax(QKᵀ)+β]V using only the public flash/sdpa interface — no kernel modification required. Derivation ---------- The paper formula expands distributively: [α · softmax(QKᵀ/√dk) + β] · V = α · [softmax(QKᵀ/√dk) · V] ← term 1: standard backend output + β · [Σ_{j∈A(i)} V_j] ← term 2: same valid keys A(i) For global causal attention, A(i) is the prefix j≤i and term 2 is a cumsum. For sliding-window or packed/additive-mask paths, term 2 must use the corresponding windowed or explicit valid-key sum. Dropout ------- The eager path drops entries of the combined weight matrix (α·softmax + β) before multiplying by V. With the flash interface we cannot access that combined matrix, so we apply dropout=0 to the flash kernel and instead apply nn.functional.dropout to the final combined output tensor. This is output dropout rather than weight dropout — a different (but standard) regularisation that achieves the same intent without the intermediate weight matrix. During inference dropout=0 so the paths are identical. Mask handling ------------- The β value-sum must respect the same valid keys as the backend attention. A 2D padding mask zeros invalid keys before prefix/window sums. A 4D additive mask is used directly to compute Σ_{j∈A(i)} V_j exactly. Memory overhead vs standard flash call --------------------------------------- Global causal attention adds one [B, H_q, S, d_head] value-sum tensor. Exact 4D-mask fallback may materialise a [B, H_q, S_q, S_k] boolean mask. Args: alpha: [B, H_q, S_q, 1] — input-dependent scale per query, in [0, 1] beta: [B, H_q, S_q, 1] — moving-average bias per query """ # ── Term 1: standard flash / sdpa output ───────────────────────────── # dropout=0.0: we apply dropout to the combined output below instead. attn_impl = module.config._attn_implementation attn_fn = ( neollm_flash_attention_forward if attn_impl in {"flash_attention_2", "flash_attention_3"} else ALL_ATTENTION_FUNCTIONS[attn_impl] ) flash_out, _ = attn_fn( module, query, key, value, attention_mask, dropout=0.0, scaling=scaling, **kwargs, ) # flash_out: [B, S, H_q, d_head] — HF wrappers all return this layout # ── Term 2: β · Σ_{j∈A(i)} V_j ────────────────────────────────────── # Fast path: compute the valid-key value sum in KV-head space, then # broadcast it to query heads during the final affine combine. Since the # sum is linear, WindowSum(repeat_kv(V)) == repeat_kv(WindowSum(V)); this # avoids doing cumsum/window arithmetic on the expanded H_q heads. alpha_t = alpha.permute(0, 2, 1, 3) # [B, S_q, H_q, 1] beta_t = beta.permute(0, 2, 1, 3) # [B, S_q, H_q, 1] if attention_mask is not None and attention_mask.ndim == 4: # Arbitrary 4D additive masks may be expressed in query-head space, so # keep the exact fallback on repeated V. This path is not the regular # causal/window hot path. value_expanded = repeat_kv(value, module.num_key_value_groups) value_sum = _affine_valid_value_sum( value_expanded, attention_mask=attention_mask, query_len=query.shape[-2], sliding_window=kwargs.get("sliding_window", None), ) v_sum_t = value_sum.transpose(1, 2) # [B,S_q,H_q,d] output = alpha_t * flash_out + beta_t * v_sum_t else: # Keep the Affine β branch as a causal prefix summary even when the # attention backend itself uses a local sliding window. This restores # the original IHA/local behavior: local logits plus a cheap global # causal value context. Padding masks are still handled inside # _affine_valid_value_sum; explicit 4D masks above remain exact. value_sum_kv = _affine_valid_value_sum( value, attention_mask=attention_mask, query_len=query.shape[-2], sliding_window=None, ) # [B,H_kv,S_q,d] v_sum_t = value_sum_kv.transpose(1, 2) # [B,S_q,H_kv,d] H_q = flash_out.shape[2] H_kv = v_sum_t.shape[2] groups = module.num_key_value_groups if groups == 1 or H_q == H_kv: output = alpha_t * flash_out + beta_t * v_sum_t else: Bsz, Sq, _, Dh = flash_out.shape flash_g = flash_out.reshape(Bsz, Sq, H_kv, groups, Dh) alpha_g = alpha_t.reshape(Bsz, Sq, H_kv, groups, 1) beta_g = beta_t.reshape(Bsz, Sq, H_kv, groups, 1) output = (alpha_g * flash_g + beta_g * v_sum_t.unsqueeze(3)).reshape( Bsz, Sq, H_q, Dh ) # ── Combine and apply dropout to the full affine output ─────────────── output = nn.functional.dropout(output, p=dropout, training=module.training) # attn_weights is None — flash never exposes the softmax weight matrix. return output, None class HadamardOProj(nn.Module): """ Parameter-free Walsh–Hadamard output projection with learnable affine rescaling. Replaces the dense W_O ∈ R^{d×d} in multi-head attention with a fixed orthogonal Walsh–Hadamard Transform followed by a per-channel learnable affine: output = α ⊙ FWHT(x) + β Motivation (Aggarwal & Kumar, 2026, arXiv:2603.08343): The standard dense o_proj develops extreme condition numbers during training (κ up to 10^5 observed in practice) because the optimiser has no incentive to keep singular values balanced — some directions are amplified while others collapse toward zero. This makes the layer hostile to FP8 quantisation, which uses a single per-tensor scale and therefore loses the low-magnitude directions entirely. The Walsh–Hadamard Transform is a fixed orthogonal matrix whose singular values are all identically 1, making κ = 1 by construction. It cannot develop condition-number pathology because it has no parameters. The learnable α/β restore per-channel expressivity at a cost of 2·d parameters instead of d². Properties: - Condition number: κ = 1 (exact, permanent, by construction) - Parameters: 2·d vs d² for dense (~25% attention params saved) - Forward FLOPs: O(d log d) vs O(d²) for dense - Norm preservation: FWHT is isometric — ‖FWHT(x)‖₂ = ‖x‖₂ - FP8 friendliness: single per-tensor scale covers all directions equally - Requires: d must be a power of 2 The FWHT is implemented as an in-place iterative butterfly (Cooley-Tukey pattern over additions/subtractions) followed by 1/√d normalisation to produce an orthonormal transform (H^T H = I). No external dependency. Reference: Aggarwal, S. & Kumar, L. (2026). "Rethinking Attention Output Projection: Structured Hadamard Transforms for Efficient Transformers." arXiv:2603.08343. """ def __init__(self, dim: int, bias: bool = True): super().__init__() assert dim > 0 and (dim & (dim - 1)) == 0, ( f"HadamardOProj requires dim to be a power of 2, got {dim}" ) self.dim = dim self.norm = dim**-0.5 # 1/√d — makes H^T H = I # Learnable affine rescaling: α ⊙ FWHT(x) + β # Initialised to α=1, β=0 so the layer starts as a pure WHT, # identical to an orthonormal projection with unit gain. self.alpha = nn.Parameter(torch.ones(dim)) self.beta = nn.Parameter(torch.zeros(dim)) if bias else None def _fwht(self, x: torch.Tensor) -> torch.Tensor: """ Iterative in-place Fast Walsh–Hadamard Transform over the last dim. Butterfly pattern: log₂(d) stages, each pairing elements at stride h. Cost: d·log₂(d) additions/subtractions, zero multiplications. Compatible with torch.compile — all shapes are static, no Python loops visible to the tracer once d is fixed. """ h = 1 while h < self.dim: # Reshape to expose pairs at current stride x = x.reshape(*x.shape[:-1], -1, 2 * h) a, b = x[..., :h], x[..., h:] # Butterfly: (a+b, a-b) — only additions and subtractions x = torch.cat([a + b, a - b], dim=-1) x = x.reshape(*x.shape[:-2], self.dim) h *= 2 return x def forward( self, x: torch.Tensor, ) -> torch.Tensor: """ Args: x: [..., dim] — concatenated multi-head attention outputs Returns: α ⊙ (FWHT(x) / √dim) + β of shape [..., dim] """ # Keep compile-time scalar metadata on the tensor device. This remains # a fused scalar multiply and allocates only one FP32 value. norm = torch.as_tensor(self.norm, device=x.device, dtype=torch.float32) out = self._fwht(x) * norm # normalise: H^T H = I out = out * self.alpha # per-channel learnable scale if self.beta is not None: out = out + self.beta # per-channel learnable bias return out class REPOModule(nn.Module): """ Context Re-Positioning module f_ϕ (Li et al., 2026, arXiv:2512.14391). Replaces the fixed linear integer indices ``0…L-1`` fed to RoPE with continuous, data-dependent positions ``z_i`` learned end-to-end. Architecture (Eq. 4–6 of the paper): # Position representation — shared across all heads in this layer r_i = Swish(h_i W_g) ⊙ (h_i W_c) r_i ∈ R^{d_p} # Position assignment — independent per head z_i^(h) = r_i w_z^(h) z_i^(h) ∈ R (scalar) where ``h_i ∈ R^d`` is the hidden state of token ``i`` entering the decoder layer (pre-FANLayer), and ``d_p = hidden_size // 8`` by default. The resulting assignments ``z [B, H, S]`` are real-valued and unconstrained. In the original REPO-RoPE path they are used directly to compute per-head ``cos/sin`` embeddings. When REPO is composed with REPO-GRAPE, ``RepoGrapePositioning`` keeps ``z`` as the raw REPO output and constructs the final coordinate u_i^(h) = i + alpha_h (z_i^(h) - i) before the GRAPE phase is formed. This keeps the semantic distinction ``REPO -> z`` versus ``REPO-GRAPE -> u`` explicit in the implementation. Design notes: - ``W_g`` and ``W_c`` are shared across heads (parameter efficiency). - ``W_z`` is a single ``[d_p, num_heads]`` matrix; each column is the per-head assignment vector ``w_z^(h)``. Vectorized as one matmul. - The raw hidden state ``h_i`` (not the FAN-augmented or normed variant) is used as input, matching the paper's formulation and avoiding circular dependency with q/k norm. - No bias on any projection — consistent with the paper's Eq. 4–5. Reference: Li, H., Zhao, T., Cai, D. & Sproat, R. (2026). "REPO: Language Models with Context Re-Positioning." arXiv:2512.14391. """ def __init__(self, hidden_size: int, d_p: int, num_heads: int): super().__init__() self.hidden_size = hidden_size self.d_p = d_p self.num_heads = num_heads # SwiGLU position representation (shared across heads, Eq. 4) self.W_g = nn.Linear(hidden_size, d_p, bias=False) self.W_c = nn.Linear(hidden_size, d_p, bias=False) # Per-head position assignment (vectorized, Eq. 5) # W_z[:, h] is w_z^(h) for head h self.W_z = nn.Linear(d_p, num_heads, bias=False) def forward( self, hidden_states: torch.Tensor, ) -> torch.Tensor: """ Args: hidden_states: [B, S, hidden_size] — residual stream entering the decoder layer, before FANLayer augmentation. Returns: z: [B, H, S] — continuous per-head position scalars. z[:, h, i] is the position assigned to token i by head h. """ # Position representation (Eq. 4): Swish(h W_g) ⊙ (h W_c) r = F.silu(self.W_g(hidden_states)) * self.W_c(hidden_states) # [B, S, d_p] # Per-head assignment (Eq. 5): z^(h) = r W_z[:, h] # W_z output: [B, S, H] → transpose to [B, H, S] z = self.W_z(r).transpose(1, 2).contiguous() # [B, H, S] return z def _apply_repo_rope( q: torch.Tensor, k: torch.Tensor, z: torch.Tensor, inv_freq: torch.Tensor, attention_scaling: float, ) -> Tuple[torch.Tensor, torch.Tensor]: """ Apply RoPE to Q and K using continuous per-head positions from REPO. Replaces the standard ``apply_rotary_pos_emb(q, k, cos, sin)`` call for layers where REPO is active. Builds ``cos/sin`` inline from ``z`` and ``inv_freq`` so that the rotation is differentiable w.r.t. ``z`` and therefore w.r.t. the parameters of REPOModule. Args: q: [B, H, S, head_dim] k: [B, H_kv, S, head_dim] (GQA: H_kv ≤ H) z: [B, H_repo, S] — per-head positions from REPOModule. Under IHA(P>1), pseudo-heads inherit the continuous position of their parent query head before RoPE. inv_freq: [rotary_dim/2] — frozen RoPE frequency vector attention_scaling: float — scaling factor from NeoLLMRotaryEmbedding Returns: (q_embed, k_embed) with the same shapes as (q, k). Implementation note on GQA: Q has ``num_attention_heads`` heads; K/V have ``num_key_value_heads`` heads (fewer under GQA). REPO produces one position per Q head. For K we average the positions of the Q heads that map to each KV head (groups of size ``num_key_value_groups``). This is the minimal approach consistent with the paper's per-head independence claim: each KV head receives a position that is representative of the Q heads it serves. """ B, H_repo, S = z.shape H_q_eff = q.shape[1] H_k_eff = k.shape[1] if H_q_eff % H_repo != 0: raise ValueError( f"REPO expected q heads divisible by z heads, got " f"H_q={H_q_eff} vs H_repo={H_repo}." ) P = H_q_eff // H_repo if H_k_eff % max(P, 1) != 0: raise ValueError( f"REPO expected k heads divisible by pseudo factor P, got " f"H_k={H_k_eff} vs P={P}." ) # Keep the logical factorization (original_query_head, pseudo_slot) # explicit so both Q and K receive positions in the same IHA pseudo-slot. z_q_struct = z.unsqueeze(2).expand(B, H_repo, P, S) # [B, H_q0, P, S] H_k_base = H_k_eff // max(P, 1) if H_repo % H_k_base != 0: raise ValueError( f"REPO expected original query heads divisible by key head groups, got " f"H_repo={H_repo} vs H_k_base={H_k_base}." ) q_per_k_base = H_repo // H_k_base z_q = z_q_struct.reshape(B, H_q_eff, S) # [B, H_q_eff, S] z_k_struct = z_q_struct.view(B, H_k_base, q_per_k_base, P, S).mean(dim=2) z_k_heads = z_k_struct.reshape(B, H_k_eff, S) # [B, H_k_eff, S] H_q = z_q.shape[1] H_kv = z_k_heads.shape[1] if H_q % H_kv != 0: raise ValueError( f"REPO expected query heads divisible by key heads, got H_q={H_q} vs H_kv={H_kv}." ) rotary_dim = inv_freq.shape[0] * 2 # inv_freq covers half the rotary dim # inv_freq arrives from rotary_emb at forward time via repo_rope_args — # already float32 on the correct device, no .to() needed, no DeviceCopy op. # No autocast barrier: explicit .float() casts on z_q/z_k are sufficient # to maintain float32 precision for the trig ops. Removing the context # manager lets Inductor plan all intermediate tensors as part of a single # static memory graph, eliminating mid-forward allocations that cause # VRAM variance under max-autotune. inv_freq_f = inv_freq # z_q: [B, H_q, S, 1] × inv_freq: [rotary_dim/2] → [B, H_q, S, rotary_dim/2] z_q = z_q.float().unsqueeze(-1) # [B, H_q, S, 1] freqs_q = z_q * inv_freq_f # [B, H_q, S, r/2] emb_q = torch.cat([freqs_q, freqs_q], dim=-1) # [B, H, S, r] attention_scale_f = torch.as_tensor( attention_scaling, device=emb_q.device, dtype=emb_q.dtype, ) cos_q = (emb_q.cos() * attention_scale_f).to(q.dtype) sin_q = (emb_q.sin() * attention_scale_f).to(q.dtype) # KV positions: average the original query heads that map to each key head, # independently inside each IHA pseudo-slot. z_k = z_k_heads.float().unsqueeze(-1) # [B, H_kv, S, 1] freqs_k = z_k * inv_freq_f # [B, H_kv, S, r/2] emb_k = torch.cat([freqs_k, freqs_k], dim=-1) # [B, H_kv, S, r] cos_k = (emb_k.cos() * attention_scale_f).to(k.dtype) sin_k = (emb_k.sin() * attention_scale_f).to(k.dtype) # Rotate only the first rotary_dim channels; pass the rest through unchanged. q_rot, q_pass = q[..., :rotary_dim], q[..., rotary_dim:] k_rot, k_pass = k[..., :rotary_dim], k[..., rotary_dim:] q_embed = torch.cat( [(q_rot * cos_q) + (rotate_half(q_rot) * sin_q), q_pass], dim=-1 ) k_embed = torch.cat( [(k_rot * cos_k) + (rotate_half(k_rot) * sin_k), k_pass], dim=-1 ) return q_embed, k_embed class RepoGrapePositioning(nn.Module): r""" REPO-GRAPE-M positional operator with a well-conditioned GQA projection. Revision implemented here ------------------------- This version keeps the contextual coordinate map unchanged, z_i^(h) = f_phi^(h)(h_i), u_i^(h) = (1 - alpha_h) i + alpha_h z_i^(h) = i + alpha_h (z_i^(h) - i), and applies exactly three corrections identified by the independent REPO-GRAPE analysis: 1. Spectral gauge fixing. The raw per-head log spectrum s_{h,r} is centered across the active rotary planes before exponentiation, s~_{h,r} = s_{h,r} - mean_r s_{h,r}, theta_{h,r} = inv_freq_r * exp(s~_{h,r}). Hence sum_r s~_{h,r}=0 and the geometric mean of exp(s~) is one. The learned spectrum can reshape relative frequencies but cannot absorb a uniform scale that is functionally confounded with u. Autograd applies the matching zero-mean projection to the raw-s gradient. 2. Learned GQA head importance. If H_m is the set of query heads sharing KV head m, one FP32 logit a_h per query head defines w_h = softmax_{h in H_m}(a_h), sum_{h in H_m} w_h = 1. a_h starts at zero, so the weighted circular reduction starts exactly at the previous uniform weighting. These are structural per-head scalars, not a token-dependent router, so there is no activation-sized routing tensor. 3. Weak circular prior toward the canonical RoPE phase. Let phi_{h,i,r} = u_i^(h) * theta_{h,r}, psi0_{i,r} = i * inv_freq_r, Z_{m,i,r} = sum_{h in H_m} w_h exp(i*phi_{h,i,r}). The shared key rotation is obtained from Z_lambda = Z + lambda * exp(i*psi0), exp(i*psi*) = Z_lambda / |Z_lambda|. This is the exact minimizer of the regularized chordal objective sum_h w_h ||R(phi_h)k - R(psi)k||_2^2 + lambda ||R(psi0)k - R(psi)k||_2^2. When the query heads are coherent, |Z| >> lambda and the prior is almost invisible. When their resultant is small, the canonical RoPE direction supplies a weak reference. Exact cancellation Z_lambda=0 remains mathematically non-identifiable; at that machine-zero event the code uses the canonical RoPE direction as a deterministic finite representative. The default prior strength is 1e-2 relative to a weighted resultant whose maximum magnitude is one. It is read from the optional config attribute ``repo_grape_circular_prior_strength`` and written back to ``config`` so the effective value is serialized. Setting it to zero is the clean prior ablation; centered spectrum and learned GQA weights remain active. Runtime positions and IHA ------------------------- ``i`` is the actual runtime ``position_ids`` value. With IHA sequence expansion, both the contextual coordinate and its canonical reference use u_{i,p} = P*u_i + p, i_{i,p} = P*i + p, so the prior points to the RoPE phase of the same pseudo-slot being rotated. Storage / checkpoint contract ----------------------------- - ``position_blend_alpha``: FP32 Parameter [H_q], unchanged. - ``freq_log_scale``: FP32 Parameter [H_q, head_dim/2], unchanged in shape; its effective value is centered before use. - ``gqa_head_weight_logits``: new FP32 Parameter [H_q], initialized to zero; old checkpoints receive this zero vector during strict loading. - ``repo_grape_circular_prior_strength``: scalar config hyperparameter, not a learned Parameter and not an activation-sized buffer. - ``u``, centered s~, phases, resultants and reference phases are transient autograd intermediates only. The current optional CCE/Triton REPO-GRAPE kernel implements the preceding legacy uncentered/unweighted/unregularized contract. It therefore remains disabled in ``NeoLLMAttention`` until that kernel is updated in lockstep; silently dispatching it would change the model mathematics. """ def __init__( self, config: NeoLLMConfig, layer_idx: int, num_attention_heads: int, num_key_value_heads: int, head_dim: int, ): super().__init__() self.config = config self.layer_idx = layer_idx self.num_attention_heads = int(num_attention_heads) self.num_key_value_heads = int(num_key_value_heads) self.head_dim = int(head_dim) self.max_rot_half = self.head_dim // 2 # Linear-blend coefficient alpha_h, one scalar per query head: # # u_i^(h) = i + alpha_h * (z_i^(h) - i) # # Initialize at alpha_h=1 so the first forward is exactly the previous # REPO behavior u=z. This is a real nn.Parameter, therefore it is # optimized by autograd/AdamW and stored in every state_dict checkpoint. # It is deliberately unconstrained: no sigmoid/clamp and no new flag. self.position_blend_alpha = nn.Parameter( torch.ones(self.num_attention_heads, dtype=torch.float32), requires_grad=True, ) # Raw learned log-spectrum. The stored parameter keeps the historical # shape/state-dict key, but _query_freq fixes the common-scale gauge via # # s_tilde[h,r] = s[h,r] - mean_r(s[h,r]). # # This preserves relative per-plane spectral reshaping while removing the # uniform log-frequency mode that can trade against u in phi=u*theta. self.freq_log_scale = nn.Parameter( torch.zeros( self.num_attention_heads, self.max_rot_half, dtype=torch.float32 ), requires_grad=True, ) # Structural importance logits for the query heads inside each GQA group. # Softmax is group-local. Zero initialization gives exactly uniform # weights, so the new degree of freedom begins from the old reduction. self.gqa_head_weight_logits = nn.Parameter( torch.zeros(self.num_attention_heads, dtype=torch.float32), requires_grad=True, ) # Weak canonical circular prior. The weighted head resultant has |Z|<=1, # so lambda=1e-2 means a 1% reference vector on the same scale. Keep this # as a config hyperparameter rather than a Parameter to avoid optimizer # state while preserving run reproducibility. self.circular_prior_strength = float( getattr(config, "repo_grape_circular_prior_strength", 1.0e-2) ) if ( not math.isfinite(self.circular_prior_strength) or self.circular_prior_strength < 0.0 ): raise ValueError( "repo_grape_circular_prior_strength must be a finite non-negative scalar." ) setattr( self.config, "repo_grape_circular_prior_strength", self.circular_prior_strength, ) def _load_from_state_dict( self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs, ): # Backward-compatible checkpoint migration. # # Very old checkpoints may lack alpha_h, and every pre-correction # checkpoint lacks the new GQA importance logits. alpha_h=1 preserves # the historical u=z initialization; a_h=0 restores uniform head weights. # Centering freq_log_scale intentionally changes only its common per-head # mode and therefore needs no new state tensor. alpha_key = prefix + "position_blend_alpha" weight_key = prefix + "gqa_head_weight_logits" need_copy = alpha_key not in state_dict or weight_key not in state_dict local_state_dict = state_dict.copy() if need_copy else state_dict if alpha_key not in state_dict: local_state_dict[alpha_key] = torch.ones( self.num_attention_heads, dtype=torch.float32 ) if weight_key not in state_dict: local_state_dict[weight_key] = torch.zeros( self.num_attention_heads, dtype=torch.float32 ) super()._load_from_state_dict( local_state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs, ) def _apply(self, fn): # Keep the angular control parameters in FP32 under model.to(dtype). # Phases/resultants are evaluated in FP32, so quantizing these tiny state # tensors would save negligible memory while perturbing the geometry. super()._apply(fn) for p in ( self.position_blend_alpha, self.freq_log_scale, self.gqa_head_weight_logits, ): p.data = p.data.float() if p.grad is not None: p.grad.data = p.grad.data.float() return self def _resolve_runtime_positions( self, batch_size: int, seq_len: int, device: torch.device, position_ids: Optional[torch.Tensor], ) -> torch.Tensor: """Resolve canonical runtime positions as FP32 ``[B,S]`` metadata.""" if position_ids is None: return torch.arange( seq_len, device=device, dtype=torch.float32 ).view(1, seq_len).expand(batch_size, -1) pid = position_ids if pid.dim() == 1: pid = pid.unsqueeze(0) elif pid.dim() > 2: # Match the existing IHA convention: metadata column 0 is the # sequential position used by RoPE/REPO. pid = pid[..., 0] if pid.dim() != 2 or pid.shape[-1] != seq_len: raise ValueError( "REPO-GRAPE position_ids must have shape [S] or [B,S] " f"with S={seq_len}, got {tuple(pid.shape)}." ) if pid.shape[0] == 1 and batch_size != 1: pid = pid.expand(batch_size, -1) elif pid.shape[0] != batch_size: raise ValueError( f"REPO-GRAPE position_ids batch={pid.shape[0]} does not match " f"batch={batch_size}." ) return pid.to(device=device, dtype=torch.float32) def transform_positions( self, z: torch.Tensor, position_ids: Optional[torch.Tensor], *, return_reference: bool = False, ) -> Union[torch.Tensor, Tuple[torch.Tensor, torch.Tensor]]: r"""Build the linear-blend coordinate and optionally return its anchor. Active map: u_i^(h) = (1-alpha_h) i + alpha_h z_i^(h) = i + alpha_h (z_i^(h) - i). ``return_reference=True`` additionally returns canonical FP32 positions ``i`` as [B,S]. They are reused to construct the RoPE prior phase after any IHA pseudo-slot expansion. Standard autograd gives dL/dz_i^(h) = alpha_h * ubar_i^(h), dL/dalpha_h = sum_{b,i} ubar_{b,h,i}(z_{b,h,i}-i). """ B, H, S = z.shape if H != self.num_attention_heads: raise ValueError( f"REPO-GRAPE expected z heads={self.num_attention_heads}, got {H}." ) pid = self._resolve_runtime_positions(B, S, z.device, position_ids) i = pid.unsqueeze(1) alpha = self.position_blend_alpha.view(1, H, 1) u = i + alpha * (z.float() - i) if return_reference: return u, pid return u def _query_freq( self, inv_freq: torch.Tensor, H_repo: int, ) -> torch.Tensor: rot_half = int(inv_freq.shape[0]) if rot_half > self.max_rot_half: raise ValueError( f"REPO-GRAPE rotary_dim/2={rot_half} exceeds head_dim/2={self.max_rot_half}." ) if H_repo != self.num_attention_heads: raise ValueError( f"REPO-GRAPE expected H_repo={self.num_attention_heads}, got {H_repo}." ) base = inv_freq.float().view(1, rot_half) raw_log_scale = self.freq_log_scale[:, :rot_half].float() # Spectral gauge fixing on the active rotary planes: # # s_tilde = s - mean_r(s), # theta = theta0 * exp(s_tilde). # # The centering matrix is an orthogonal projection onto the zero-mean # subspace, so autograd also removes the common-mode gradient. centered_log_scale = raw_log_scale - raw_log_scale.mean( dim=-1, keepdim=True ) freq = base * torch.exp(centered_log_scale) return freq.to(device=inv_freq.device) def apply_multiplicative( self, q: torch.Tensor, k: torch.Tensor, u: torch.Tensor, inv_freq: torch.Tensor, attention_scaling: float, *, reference_positions: Optional[torch.Tensor] = None, ) -> Tuple[torch.Tensor, torch.Tensor]: r"""Apply corrected REPO-GRAPE-M rotations to Q/K. Query heads keep their individual phases phi_{h,i,r} = u_i^(h) * theta0_r * exp(s~_{h,r}). For each GQA group H_m, structural logits a_h define w_h = softmax_{h in H_m}(a_h), Z = sum_h w_h exp(i*phi_h). The shared key uses the weakly regularized projection psi0 = i * theta0, Z_lambda = Z + lambda exp(i*psi0), exp(i*psi*) = Z_lambda / |Z_lambda|. This is still a closed-form SO(2) projection: it minimizes the weighted chordal key distortion plus lambda times the same distortion to the canonical RoPE rotation. The prior acts only on the shared GQA key; query phases remain fully REPO-GRAPE-specific. ``reference_positions`` must describe the same final sequence axis as ``u``. The attention caller passes P*i+p when IHA expands pseudo-slots. """ B, H_repo, S = u.shape H_q_eff = q.shape[1] H_k_eff = k.shape[1] if H_q_eff % H_repo != 0: raise ValueError( f"REPO-GRAPE expected q heads divisible by coordinate heads, " f"got H_q={H_q_eff} vs H_repo={H_repo}." ) P = H_q_eff // H_repo if H_k_eff % max(P, 1) != 0: raise ValueError( f"REPO-GRAPE expected k heads divisible by pseudo factor P, " f"got H_k={H_k_eff} vs P={P}." ) u_q_struct = u.unsqueeze(2).expand(B, H_repo, P, S) # [B,H_q_base,P,S] H_k_base = H_k_eff // max(P, 1) if H_repo % H_k_base != 0: raise ValueError( f"REPO-GRAPE expected original query heads divisible by key groups, " f"got H_repo={H_repo} vs H_k_base={H_k_base}." ) q_per_k_base = H_repo // H_k_base u_q = u_q_struct.reshape(B, H_q_eff, S) rot_half = int(inv_freq.shape[0]) rotary_dim = rot_half * 2 freq_repo = self._query_freq(inv_freq, H_repo) freq_q_struct = freq_repo.unsqueeze(1).expand(H_repo, P, rot_half) freq_q = freq_q_struct.reshape(H_q_eff, rot_half) # Individual query phases are the fundamental quantity. They are # calculated in fp32 before trigonometric evaluation, as in the # original REPO-GRAPE path. phase_q = u_q.float().unsqueeze(-1) * freq_q.to(device=q.device).view( 1, H_q_eff, 1, rot_half ) cos_q_half = phase_q.cos() sin_q_half = phase_q.sin() # Corrected GQA circular projection. Restore the logical base-head # grouping and reduce only over query-within-group. Structural weights # are token-independent and initialized uniformly. cos_q_grouped = cos_q_half.reshape( B, H_k_base, q_per_k_base, P, S, rot_half ) sin_q_grouped = sin_q_half.reshape( B, H_k_base, q_per_k_base, P, S, rot_half ) head_logits = self.gqa_head_weight_logits.reshape( H_k_base, q_per_k_base ).float() head_weights = torch.softmax(head_logits, dim=-1).view( 1, H_k_base, q_per_k_base, 1, 1, 1 ) cos_k_half = (cos_q_grouped * head_weights).sum(dim=2) sin_k_half = (sin_q_grouped * head_weights).sum(dim=2) # Canonical RoPE prior for the same runtime/pseudo-slot positions. ref = self._resolve_runtime_positions(B, S, q.device, reference_positions) base_phase_half = ref.view(B, 1, 1, S, 1) * inv_freq.float().to( device=q.device ).view(1, 1, 1, 1, rot_half) base_cos_half = base_phase_half.cos() base_sin_half = base_phase_half.sin() prior_strength = torch.as_tensor( self.circular_prior_strength, device=q.device, dtype=cos_k_half.dtype, ) cos_k_half = cos_k_half + prior_strength * base_cos_half sin_k_half = sin_k_half + prior_strength * base_sin_half # Project Z_lambda back onto S^1. The weak prior regularizes the usual # low-coherence regime but cannot forbid the exact cancellation # Z+lambda*exp(i*psi0)=0. At that machine-zero event every angle is a # minimizer of the regularized objective, so use the canonical phase as # a deterministic finite representative rather than identity-at-zero. resultant_sq = cos_k_half.square() + sin_k_half.square() well_defined = resultant_sq > torch.finfo(resultant_sq.dtype).eps safe_resultant_sq = torch.where( well_defined, resultant_sq, torch.ones_like(resultant_sq) ) inv_resultant = torch.rsqrt(safe_resultant_sq) cos_k_half = torch.where( well_defined, cos_k_half * inv_resultant, base_cos_half.expand_as(cos_k_half), ) sin_k_half = torch.where( well_defined, sin_k_half * inv_resultant, base_sin_half.expand_as(sin_k_half), ) cos_k_half = cos_k_half.reshape(B, H_k_eff, S, rot_half) sin_k_half = sin_k_half.reshape(B, H_k_eff, S, rot_half) attention_scale_f = torch.as_tensor( attention_scaling, device=phase_q.device, dtype=phase_q.dtype, ) cos_q = ( torch.cat([cos_q_half, cos_q_half], dim=-1) * attention_scale_f ).to(q.dtype) sin_q = ( torch.cat([sin_q_half, sin_q_half], dim=-1) * attention_scale_f ).to(q.dtype) cos_k = ( torch.cat([cos_k_half, cos_k_half], dim=-1) * attention_scale_f ).to(k.dtype) sin_k = ( torch.cat([sin_k_half, sin_k_half], dim=-1) * attention_scale_f ).to(k.dtype) q_rot, q_pass = q[..., :rotary_dim], q[..., rotary_dim:] k_rot, k_pass = k[..., :rotary_dim], k[..., rotary_dim:] q_embed = torch.cat( [(q_rot * cos_q) + (rotate_half(q_rot) * sin_q), q_pass], dim=-1, ) k_embed = torch.cat( [(k_rot * cos_k) + (rotate_half(k_rot) * sin_k), k_pass], dim=-1, ) return q_embed, k_embed class RepoGoatPrior(nn.Module): r""" Factorised GOAT-style attention log-prior for NeoLLM. GOAT interprets an additive attention bias as a log-prior inside the row-wise KL/EOT objective: p_i = softmax(s_i / tau + log pi_i). This module keeps that term separate from the REPO-GRAPE-M geometric score. Instead of materialising K_ij as a dense [S,S] matrix, it appends a small positional subspace to Q and K: [Q | Q_prior] [K | K_prior]^T / sqrt(d_head) = QK^T / sqrt(d_head) + K_prior_logits. V is padded with zeros in the appended dimensions and the attention output is sliced back to head_dim by the caller, so downstream modules see exactly the same shape whether the prior is enabled or disabled. Components ---------- 1. Recency key-only prior: q_r = lambda_h, k_r = c_j This is equivalent to lambda_h(c_j - c_i) modulo a per-row constant. 2. Relative Fourier prior: alpha_h,r cos(omega_r(c_i-c_j)) + beta_h,r sin(omega_r(c_j-c_i)) represented by two separable coordinates per frequency. 3. Sink key-only prior: q_s = rho_h, k_s = exp(-softplus(delta_h) c_j) This gives an explicit, controllable sink profile instead of forcing the model to create sinks indirectly through content norms. Initialisation is an exact no-op: all amplitudes are zero, while the key-side bases remain non-zero so gradients immediately reach the amplitudes. Parameters are kept in fp32 under model.to(dtype), matching the REPO-GRAPE spectral parameters. """ def __init__( self, config: NeoLLMConfig, layer_idx: int, num_attention_heads: int, head_dim: int, ): super().__init__() self.config = config self.layer_idx = layer_idx self.num_attention_heads = int(num_attention_heads) self.head_dim = int(head_dim) self.num_frequencies = int(getattr(config, "repo_goat_num_frequencies", 3)) self.prior_dim = 2 * self.num_frequencies + 2 self.sqrt_head_dim = math.sqrt(float(self.head_dim)) # Amplitudes live on the query side so each query head can learn a # different prior even under GQA, where K/V heads are shared by groups. self.recency_slope = nn.Parameter( torch.zeros(self.num_attention_heads, dtype=torch.float32) ) self.rel_alpha = nn.Parameter( torch.zeros( self.num_attention_heads, self.num_frequencies, dtype=torch.float32 ) ) self.rel_beta = nn.Parameter( torch.zeros( self.num_attention_heads, self.num_frequencies, dtype=torch.float32 ) ) self.sink_strength = nn.Parameter( torch.zeros(self.num_attention_heads, dtype=torch.float32) ) # Learnable positive sink decay, initialized to the configured value. sink_decay = float(getattr(config, "repo_goat_sink_decay", 4.0)) inv_softplus = math.log(math.exp(sink_decay) - 1.0) self.sink_log_decay = nn.Parameter( torch.full((self.num_attention_heads,), inv_softplus, dtype=torch.float32) ) if self.num_frequencies > 0: base_freq = ( 2.0 * math.pi * torch.arange(1, self.num_frequencies + 1, dtype=torch.float32) ) else: base_freq = torch.empty(0, dtype=torch.float32) self.register_buffer("prior_freq", base_freq, persistent=False) def _apply(self, fn): # Keep small prior amplitudes in fp32 under model.to(dtype). super()._apply(fn) for p in ( self.recency_slope, self.rel_alpha, self.rel_beta, self.sink_strength, self.sink_log_decay, ): p.data = p.data.float() if p.grad is not None: p.grad.data = p.grad.data.float() self.prior_freq = self.prior_freq.float() return self def _normalised_positions( self, seq_len: int, batch_size: int, device: torch.device, position_ids: Optional[torch.Tensor], ) -> torch.Tensor: """ Return a differentiable-free causal coordinate c in [0,1]. If position_ids are supplied, they are respected, including the IHA seq-expanded case where the attention sequence length is S*P but position_ids still describe the original S tokens. When no compatible position_ids exist, fall back to arange(seq_len). """ pos = None if position_ids is not None: pid = position_ids if pid.dim() == 1: pid = pid.unsqueeze(0) elif pid.dim() > 2: pid = pid[..., 0] base_len = pid.shape[-1] if base_len == seq_len: pos = pid elif base_len > 0 and seq_len % base_len == 0: P = seq_len // base_len offsets = torch.arange(P, device=pid.device, dtype=pid.dtype) pos = (pid.unsqueeze(-1) * P + offsets.view(1, 1, P)).reshape( pid.shape[0], seq_len ) if pos is None: pos = torch.arange(seq_len, device=device, dtype=torch.float32).view( 1, seq_len ) else: pos = pos.to(device=device, dtype=torch.float32) if pos.shape[0] == 1 and batch_size != 1: pos = pos.expand(batch_size, -1) elif pos.shape[0] != batch_size: pos = pos[:1].expand(batch_size, -1) pos = pos - pos.amin(dim=-1, keepdim=True) denom = pos.amax(dim=-1, keepdim=True).clamp_min(1.0) return (pos / denom).clamp(0.0, 1.0) def append_prior_subspace( self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, position_ids: Optional[torch.Tensor] = None, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, int]: """ Append prior dimensions to q/k/v and return the appended width. Args: q: [B, H_q, S_q, head_dim] k: [B, H_kv, S_k, head_dim] v: [B, H_kv, S_k, head_dim] Returns: q_ext, k_ext, v_ext, prior_dim """ B, H_q, S_q, _ = q.shape H_kv, S_k = k.shape[1], k.shape[2] if H_q != self.num_attention_heads: raise ValueError( f"REPO-GOAT expected H_q={self.num_attention_heads}, got {H_q}." ) c_q = self._normalised_positions(S_q, B, q.device, position_ids) c_k = self._normalised_positions(S_k, B, k.device, position_ids) q_prior = q.new_zeros(B, H_q, S_q, self.prior_dim) k_prior = k.new_zeros(B, H_kv, S_k, self.prior_dim) c_q_f = c_q.float().unsqueeze(1) # [B,1,S_q] c_k_f = c_k.float().unsqueeze(1) # [B,1,S_k] scale = self.sqrt_head_dim # 0: recency. softmax ignores the missing row-constant -lambda_h c_i, # so lambda_h*c_j is distributionally equivalent to lambda_h(c_j-c_i). q_prior[..., 0] = (self.recency_slope.float().view(1, H_q, 1) * scale).to( q.dtype ) k_prior[..., 0] = c_k_f.expand(B, H_kv, S_k).to(k.dtype) # 1..2R: relative Fourier kernel. Coefficients are query-head specific; # key-side features are shared over KV heads, preserving GQA semantics. if self.num_frequencies > 0: omega = self.prior_freq.to(device=q.device, dtype=torch.float32).view( 1, 1, 1, -1 ) angle_q = c_q_f.unsqueeze(-1) * omega # [B,1,S_q,R] angle_k = c_k_f.unsqueeze(-1) * omega # [B,1,S_k,R] cos_q = angle_q.cos() sin_q = angle_q.sin() cos_k = angle_k.cos() sin_k = angle_k.sin() alpha = self.rel_alpha.float().view(1, H_q, 1, self.num_frequencies) beta = self.rel_beta.float().view(1, H_q, 1, self.num_frequencies) q_cos = (alpha * cos_q - beta * sin_q) * scale q_sin = (alpha * sin_q + beta * cos_q) * scale q_prior[..., 1 : 1 + 2 * self.num_frequencies : 2] = q_cos.to(q.dtype) q_prior[..., 2 : 1 + 2 * self.num_frequencies : 2] = q_sin.to(q.dtype) k_prior[..., 1 : 1 + 2 * self.num_frequencies : 2] = cos_k.expand( B, H_kv, S_k, self.num_frequencies ).to(k.dtype) k_prior[..., 2 : 1 + 2 * self.num_frequencies : 2] = sin_k.expand( B, H_kv, S_k, self.num_frequencies ).to(k.dtype) # Last coordinate: key-only sink profile. sink_profile_k = torch.exp( -F.softplus(self.sink_log_decay.float()).mean() * c_k_f ) q_prior[..., -1] = (self.sink_strength.float().view(1, H_q, 1) * scale).to( q.dtype ) k_prior[..., -1] = sink_profile_k.expand(B, H_kv, S_k).to(k.dtype) # Value prior channels are zeros: prior affects weights, never values. v_prior = v.new_zeros(B, v.shape[1], v.shape[2], self.prior_dim) q_ext = torch.cat([q, q_prior], dim=-1).contiguous() k_ext = torch.cat([k, k_prior], dim=-1).contiguous() v_ext = torch.cat([v, v_prior], dim=-1).contiguous() return q_ext, k_ext, v_ext, self.prior_dim class NeoLLMAttention(nn.Module): """ Full attention with FANformer, RMSNorm, ResFormer, Learnable Multipliers, optional Momentum, MEA head-level composition, optional LUCID preconditioning, optional Affine-Scaled Attention, optional Exclusive Self Attention, optional Directional Routing (Taylor, 2026), optional Context Re-Positioning (Li et al., 2026), optional REPO-GRAPE contextual group positioning (Li et al., 2026 + Zhang et al., 2026), optional Wall-style per-plane REPO-GRAPE retention, and optional REPO-GOAT factorised log-priors (Litman & Guo, 2026). Directional Routing inserts at position C — post-XSA, pre-reshape — where the output is already normalized (MEAHeadRMSNorm) and has auto-position removed (XSA). The router suppresses directions of cross-domain interference orthogonal to the self-position already cleaned by XSA. Pipeline (all active simultaneously when enabled): FANLayer → q_proj(gate) → q_norm/k_norm → REPO/RoPE or REPO-GRAPE-M → Momentum → MEA(K,V) → LUCID(V) → v_ref → optional GOAT log-prior → optional HOLA memory-augmented FA2 read (IHA seq-expand only) → Affine-Scaled SDPA → MEAHeadRMSNorm → XSA → Directional Routing → reshape → o_proj · sigmoid(gate) → dropout RoPE variants (controlled by config.use_repo and layer_idx): use_repo=False (default): standard integer RoPE via pre-computed position_embeddings — identical to prior behaviour. use_repo=True, layer_idx >= repo_start_layer: REPOModule f_ϕ predicts continuous per-head positions z [B, H, S] from hidden_states. cos/sin are built inline from z and inv_freq so the rotation is differentiable w.r.t. f_ϕ parameters. use_repo=True, layer_idx < repo_start_layer: standard integer RoPE (lower layers capture surface features that benefit less from re-positioning). o_proj variants (controlled by config.use_hadamard_o_proj): False (default): dense LinearWithMultipliers — full expressivity, with Learnable Multipliers controlled by config.use_learnable_multipliers; develops high κ during training (FP8 risk). True: HadamardOProj — fixed WHT + learnable α/β, κ = 1 by construction, 25% fewer attention params, FP8-friendly (Aggarwal & Kumar, 2026, arXiv:2603.08343). References: Directional Routing: Taylor (2026). arXiv:2603.14923. Hadamard o_proj: Aggarwal & Kumar (2026). arXiv:2603.08343. Context Re-Positioning: Li et al. (2026). arXiv:2512.14391. """ def __init__(self, config: NeoLLMConfig, layer_idx: int): super().__init__() self.config = config self.layer_idx = layer_idx self.head_dim = getattr( config, "head_dim", config.hidden_size // config.num_attention_heads ) self.num_key_value_groups = ( config.num_attention_heads // config.num_key_value_heads ) self.scaling = self.head_dim**-0.5 self.sqrt_head_dim = math.sqrt(self.head_dim) self.attention_dropout = config.attention_dropout self.is_causal = True self.use_momentum_attention = getattr(config, "use_momentum_attention", False) self.momentum_gamma = float(getattr(config, "momentum_gamma", 0.0)) self.use_mea_attention = getattr(config, "use_mea_attention", False) self.mea_component_key_value_heads = int( getattr(config, "mea_component_key_value_heads", config.num_key_value_heads) ) self.mea_groupnorm_eps = float( getattr(config, "mea_groupnorm_eps", config.rms_norm_eps) ) self.use_lucid_attention = getattr(config, "use_lucid_attention", False) self.lucid_attention_eps = float( getattr(config, "lucid_attention_eps", config.rms_norm_eps) ) self.use_hadamard_o_proj = getattr(config, "use_hadamard_o_proj", False) self.use_learnable_multipliers = getattr( config, "use_learnable_multipliers", True ) self.fan_layer = FANLayer( hidden_size=config.hidden_size, fan_ratio=getattr(config, "fan_ratio", 0.125), ) fan_output_dim = config.hidden_size + int( config.hidden_size * getattr(config, "fan_ratio", 0.125) ) self.q_proj = LinearWithMultipliers( fan_output_dim, config.num_attention_heads * self.head_dim * 2, bias=config.attention_bias, use_row_multiplier=True, use_column_multiplier=False, enable_multipliers=self.use_learnable_multipliers, ) self.num_mea_component_heads = ( self.mea_component_key_value_heads if self.use_mea_attention else config.num_key_value_heads ) self.k_proj = nn.Linear( fan_output_dim, self.num_mea_component_heads * self.head_dim, bias=config.attention_bias, ) self.v_proj = nn.Linear( fan_output_dim, self.num_mea_component_heads * self.head_dim, bias=config.attention_bias, ) # ── Output projection (Aggarwal & Kumar, 2026, arXiv:2603.08343) ──── # use_hadamard_o_proj=False (default): dense LinearWithMultipliers. # use_hadamard_o_proj=True: HadamardOProj — fixed WHT + learnable α/β. # κ = 1 by construction, 25% fewer attention params, FP8-friendly. # Requires hidden_size to be a power of 2 (512 ✓, 1024 ✓, 768 ✗). _o_in = config.num_attention_heads * self.head_dim if self.use_hadamard_o_proj: assert _o_in == config.hidden_size, ( f"HadamardOProj requires in_dim == out_dim, " f"got {_o_in} vs {config.hidden_size}" ) self.o_proj = HadamardOProj(config.hidden_size, bias=config.attention_bias) else: self.o_proj = LinearWithMultipliers( _o_in, config.hidden_size, bias=config.attention_bias, use_row_multiplier=True, use_column_multiplier=True, enable_multipliers=self.use_learnable_multipliers, ) self.q_norm = _make_norm(self.head_dim, eps=config.rms_norm_eps) self.k_norm = _make_norm(self.head_dim, eps=config.rms_norm_eps) if self.use_mea_attention: self.mea_key_mix = nn.Parameter( torch.eye(self.num_mea_component_heads, config.num_key_value_heads) ) self.mea_value_mix = nn.Parameter( torch.eye(self.num_mea_component_heads, config.num_key_value_heads) ) self.mea_output_norm = MEAHeadRMSNorm( num_heads=config.num_attention_heads, head_dim=self.head_dim, num_kv_groups=self.num_key_value_groups, eps=self.mea_groupnorm_eps, ) else: self.mea_key_mix = None self.mea_value_mix = None self.mea_output_norm = None self.dropout = nn.Dropout(config.dropout_rate) self.use_fan_residual = getattr(config, "use_fan_residual", True) if self.use_fan_residual: self.lambda_1 = nn.Parameter(torch.tensor(0.5)) self.lambda_2 = nn.Parameter(torch.tensor(0.5)) else: self.lambda_1 = None self.lambda_2 = None # ── Affine-Scaled Attention (Bae et al., 2026) ─────────────────────── self.use_affine_scaled_attention = getattr( config, "use_affine_scaled_attention", False ) self.affine_momentum = float(getattr(config, "affine_momentum", 0.9)) # ── Exclusive Self Attention (Zhai, 2026) ──────────────────────────── self.use_xsa = getattr(config, "use_xsa", False) self.xsa_eps = float(getattr(config, "xsa_eps", 1e-6)) # ── HOLA-derived episodic memory (Cui, 2026) ───────────────────────── # This flag is deliberately scoped to the direct FlashAttention-2 # seq-expand branch used by paper-correct IHA (P>1). It does not add a # second attention output. Instead, each chunk performs a single FA2 call # over [selected episodic KV, current chunk KV]. Memory logits live in an # augmented subspace so ordinary keys and memory keys can use different # read geometries while still sharing one softmax. # # Exact HOLA writes use beta*||e|| from a recurrent delta state. This # Transformer branch has no such state and stock FA2 does not expose the # pre-dropout accumulators needed for exact leave-one-out residuals. # Therefore the FA2 path below uses a compile-friendly value-innovation # proxy for episode eviction. The read path itself is exact within the # single augmented FA2 softmax. self.use_hola_memory = bool(getattr(config, "use_hola_memory", False)) self.hola_memory_size = int(getattr(config, "hola_memory_size", 64)) self.hola_chunk_size = int(getattr(config, "hola_chunk_size", 256)) self.hola_score_eps = float(getattr(config, "hola_score_eps", 1e-6)) if self.use_hola_memory: self.hola_q_gamma = nn.Parameter( torch.full( (config.num_attention_heads, self.head_dim), float(getattr(config, "hola_rms_gamma_init", 1.0)), ) ) self.hola_k_gamma = nn.Parameter( torch.full( (config.num_key_value_heads, self.head_dim), float(getattr(config, "hola_rms_gamma_init", 1.0)), ) ) # Tied at KV-head granularity because memory keys are stored once per # KV component. Query heads that share a KV component share this prior. self.hola_logit_bias = nn.Parameter( torch.full( (config.num_key_value_heads,), float(getattr(config, "hola_logit_bias_init", -2.0)), ) ) else: self.hola_q_gamma = None self.hola_k_gamma = None self.hola_logit_bias = None if self.use_affine_scaled_attention: self.alpha_proj = nn.Linear( config.hidden_size, config.num_attention_heads, bias=False ) self.register_buffer( "alpha_ma", torch.zeros(1, config.num_attention_heads, 1, 1), persistent=True, ) # ── Directional Routing (Taylor, 2026) ─────────────────────────────── # Each attention head learns K unit-norm direction vectors in head-space. # A shared 4-layer MLP router — conditioned on the mean-pooled sequence # representation — produces per-input sigmoid weights r_{h,k} ∈ [0,1] # that control how much of each directional component is suppressed from # the head's output after XSA (position C): # # o'_h = o_h - Σ_k r_{h,k} · (o_h · d_{h,k}) · d_{h,k} # # Position C (post-XSA, pre-reshape) is chosen because: # - XSA already removed auto-position (self-position noise). # - MEAHeadRMSNorm already normalized the output — suppression has # predictable magnitude since d_{h,k} is unit-norm. # - Directions live in head-space d_head before o_proj, preserving # the vocabulary projection interpretability of the paper. # - No interaction with the Gated Attention gate (applied post o_proj). # # When use_xsa=False, position C reduces to post-RMSNorm, pre-reshape — # routing still applies correctly, directions just also span the # self-position subspace (no XSA cleaned it first). # # Router: mean-pools hidden_states (pre-FAN residual stream) over S, # passes through 4-layer MLP, outputs H×K logits, temperature-scaled # sigmoid → r_{h,k}. Temperature T=5.0 pushes weights toward {0,1}. # The router is shared across all heads within this layer, exactly as # in the paper. No auxiliary loss — learns from LM objective only. # # direction_vecs: [H, K, d_head] — unit-normalized in forward, not init. # Initialized from normal(0, 1) and normalized at first forward pass. self.use_directional_routing = getattr(config, "use_directional_routing", False) self.directional_routing_k = int(getattr(config, "directional_routing_k", 4)) self.directional_routing_temp = float( getattr(config, "directional_routing_temp", 5.0) ) if self.use_directional_routing: H = config.num_attention_heads K = self.directional_routing_k D = self.head_dim R = config.hidden_size # router hidden dim # Direction vectors: [H, K, d_head] # Stored unnormalized — unit-normalized during forward. self.direction_vecs = nn.Parameter(torch.randn(H, K, D)) # 4-layer MLP router shared across heads within this layer. # Input: mean-pooled hidden_states [B, hidden_size] # Output: [B, H*K] → reshape [B, H, K] → sigmoid(T·x) → r_{h,k} # Intermediate dim = hidden_size throughout, matching the paper. self.direction_router = nn.Sequential( nn.RMSNorm(R, eps=config.rms_norm_eps), nn.Linear(R, R, bias=True), nn.GELU(), nn.Linear(R, R, bias=True), nn.GELU(), nn.Linear(R, R, bias=True), nn.GELU(), nn.Linear(R, H * K, bias=True), ) else: self.direction_vecs = None self.direction_router = None # ── Context Re-Positioning / REPO-GRAPE ───────────────────────────── # REPO-GRAPE reuses the REPO coordinate module f_phi and replaces only # the positional action applied to Q/K. Enabling use_repo_grape=True # therefore activates the REPO coordinate path even when use_repo=False. # There is no extra position-mode flag: active REPO-GRAPE always maps # raw z to u=i+alpha_h*(z-i) with one learned FP32 alpha_h per query # head before GRAPE-M (and before the IHA P*u+p expansion). _repo_requested = bool(getattr(config, "use_repo", False)) _repo_grape_requested = bool(getattr(config, "use_repo_grape", False)) _repo_start_layer = getattr( config, "repo_start_layer", config.num_hidden_layers // 3 ) self.use_repo = ( _repo_requested or _repo_grape_requested ) and layer_idx >= _repo_start_layer self.use_repo_grape = _repo_grape_requested and self.use_repo if self.use_repo: _d_p = getattr(config, "repo_d_p", config.hidden_size // 8) self.repo_module = REPOModule( hidden_size=config.hidden_size, d_p=_d_p, num_heads=config.num_attention_heads, ) else: self.repo_module = None if self.use_repo_grape: self.repo_grape = RepoGrapePositioning( config=config, layer_idx=layer_idx, num_attention_heads=config.num_attention_heads, num_key_value_heads=config.num_key_value_heads, head_dim=self.head_dim, ) # The currently available CCE/Triton REPO-GRAPE kernel implements # the legacy uncentered + uniform + no-prior projection. The model # now centers the spectrum, learns structural GQA weights and adds a # weak canonical circular prior, so that kernel MUST stay disabled # until its forward/backward contract is updated in lockstep. # Otherwise Torch and Triton would silently implement different math. self.use_repo_grape_triton = False else: self.repo_grape = None self.use_repo_grape_triton = False # ── REPO-GOAT factorised log-prior ─────────────────────────────────── # Independent of use_repo/use_repo_grape: it can be ablated with standard # RoPE or composed with REPO-GRAPE-M. It activates only from # repo_start_layer onward to mirror the REPO positional stack schedule. self.use_repo_goat_prior = ( bool(getattr(config, "use_repo_goat_prior", False)) and layer_idx >= _repo_start_layer ) if self.use_repo_goat_prior: self.repo_goat_prior = RepoGoatPrior( config=config, layer_idx=layer_idx, num_attention_heads=config.num_attention_heads, head_dim=self.head_dim, ) else: self.repo_goat_prior = None # ── Interleaved Head Attention (Duvvuri et al., 2026) ───────────────── # Implementa cross-head mixing antes de RoPE: cada pseudo-head h,p # es combinación lineal aprendida de los H input-heads originales. # P=1: solo mixing (misma forma), compatible con todos los flags. # P>1: expande la secuencia real S → S·P y colapsa vía R después # de SDPA. num_key_value_groups se preserva porque Q y KV se # expanden en la dimensión de secuencia, no en la de heads. # # Schedule local-global: # 'L' → capa local con IHA y ventana W=N/(2P²). # 'G' → capa global de transporte. Por defecto NO usa IHA, porque # el paper empareja FLOPs como 4 locales IHA + 1 global estándar. # Activa config.iha_global_layers_use_iha=True solo si quieres # la ablación más cara donde la G también ejecuta global IHA. # # Si una capa 'G' no usa IHA, no se instancian iha_alpha_q/k/v ni R # en esa capa. Así el flag desactiva compute, parámetros y grafo. # Init identidad: IHA ≡ MHA en paso 0 (Teorema 2, inclusión M⊆P_P). # Referencia: Duvvuri et al. (2026), arXiv:2602.21371. self.iha_P = int(getattr(config, "iha_num_pseudo_heads", 1)) # Keep IHA topology and FlashAttention integer arguments static. A # window larger than the actual sequence is equivalent to full causal # attention, so the configured context is sufficient. self.iha_context_len = max( 1, int(getattr(config, "max_position_embeddings", 1)) ) _base_use_iha = bool(getattr(config, "use_iha", False)) _pattern = list(getattr(config, "iha_local_global_pattern", "LLLLG").upper()) _pos = layer_idx % max(len(_pattern), 1) self.iha_schedule_token = _pattern[_pos] if _pattern else "G" self.iha_global_layers_use_iha = bool( getattr(config, "iha_global_layers_use_iha", False) ) # Per-layer activation: # local slots always use IHA when global use_iha=True; # global slots use IHA only under the explicit ablation flag. self.iha_layer_uses_iha = bool( _base_use_iha and ( self.iha_schedule_token == "L" or (self.iha_schedule_token == "G" and self.iha_global_layers_use_iha) ) ) self.use_iha = self.iha_layer_uses_iha if self.use_iha: _H_q = config.num_attention_heads _H_kv = self.num_mea_component_heads _P = self.iha_P # alpha_q[h_out, h_in, p]: mezcla de h_in sobre todos los heads # → pseudo p de head h_out. # En K/V usamos H_comp (pre-MEA), de modo que MEA pueda recomponer # H_comp → H_kv independientemente dentro de cada pseudo-slot. self.iha_alpha_q = nn.Parameter(torch.zeros(_H_q, _H_q, _P)) self.iha_alpha_k = nn.Parameter(torch.zeros(_H_kv, _H_kv, _P)) self.iha_alpha_v = nn.Parameter(torch.zeros(_H_kv, _H_kv, _P)) # R[h, p]: colapsa el slot p del pseudo-token hacia head h. # Shape [H_q, P] — colapso sobre la dimensión de pseudo-slots en # la secuencia expandida (paper Alg. 1 Step 8: einsum 'hp,hnpd→hnd'). # Init identidad: slot p=0 pasa íntegro, resto cero → M⊆P_P (Teo. 2). self.iha_R = nn.Parameter(torch.zeros(_H_q, _P)) # Init identidad: IHA ≡ MHA en paso 0 (Teorema 2, inclusión M⊆P_P). with torch.no_grad(): for _h in range(_H_q): self.iha_alpha_q.data[_h, _h, 0] = 1.0 self.iha_R.data[_h, 0] = 1.0 # slot 0 → identidad for _h in range(_H_kv): self.iha_alpha_k.data[_h, _h, 0] = 1.0 self.iha_alpha_v.data[_h, _h, 0] = 1.0 # ── Schedule local-global (paper §5.1 / Appendix C) ────────────── # El pattern "LLLLG" repite cada len(pattern) capas: # 'L' → capa local IHA con sliding window (FLOP-cheap). # 'G' → capa global. Con iha_global_layers_use_iha=False esta # rama no llega aquí, porque la capa G no instancia IHA. # Con True, la G sí llega aquí y usa IHA global full attention. # Solo relevante cuando P>1; con P=1 no hay expansión. self.iha_is_local = self.iha_P > 1 and self.iha_schedule_token == "L" # Tamaño de ventana W para capas locales. # Si el usuario fija iha_sliding_window, se usa ese valor. # Si no, se fija en construcción como W = N/(2P²) usando el # contexto configurado; la longitud efectiva/padding no cambia # la topología compilada ni la formulación del paper. _cfg_window = getattr(config, "iha_sliding_window", None) self.iha_window = int(_cfg_window) if _cfg_window is not None else None # ── Flash-attn instance references (P>1 only) ───────────────── # Stored as plain attributes (not nn.Parameters/buffers) so # Dynamo sees them as compile-time constants — no global mutation # inside forward, no graph break. # Raise at __init__ time (not forward time) so the error surfaces # immediately when the model is constructed, not at the first step. if self.iha_P > 1: if not _IHA_FLASH_ATTN_AVAILABLE: raise ImportError( "Interleaved Head Attention with P>1 (seq-expand / paper-correct " "mode) requires flash_attn to be explicitly installed.\n" "Install it with:\n" " pip install flash-attn --no-build-isolation\n\n" "If flash_attn is not available, set iha_num_pseudo_heads=1 " "(P=1) to fall back to the head-expand approximation, which " "uses the standard HF attention backend." ) self._iha_flash_attn_func = _IHA_FA_FUNC self._iha_flash_attn_varlen = _IHA_FA_VARLEN else: self._iha_flash_attn_func = None self._iha_flash_attn_varlen = None else: self.iha_is_local = False self.iha_window = None self._iha_flash_attn_func = None self._iha_flash_attn_varlen = None def _resolve_iha_window(self, seq_len: Optional[int] = None) -> int: """ Resolve the local IHA sliding window without making the operator depend on the current prompt/padding length. Exact paper recipe (Sec. 5.1 / Appendix C): W := N / (2 P^2) Here N is treated as the configured context/training length unless the user provides an explicit iha_sliding_window. This keeps prefix-only, prefix+suffix, and padded forwards on the same local-attention operator. """ if self.iha_window is not None: return max(1, int(self.iha_window)) denom = max(2 * self.iha_P * self.iha_P, 1) return max(1, self.iha_context_len // denom) def _build_iha_local_mask( self, seq_len: int, window_size: int, device: torch.device, dtype: torch.dtype, ) -> torch.Tensor: """ Construye máscara aditiva de ventana deslizante [1, 1, S, S]. Implementa la restricción de la capa IHA local del paper (§5.1): cada token i solo atiende a los últimos `window_size` tokens previos (i − W ≤ j ≤ i). Se SUMA a la máscara causal existente, combinando causalidad + restricción de banda en un único tensor. Compatibilidad con backends (paper §2, FlashAttention §A): - eager : se suma a attn_weights antes del softmax → -inf → exp=0 - sdpa : misma interfaz que eager vía scaled_dot_product_attention - flash2/3: la localidad se expresa con `sliding_window` del backend; este helper solo se usa en rutas que sí consumen una máscara aditiva 4D directamente. Window auto follows the exact paper schedule: W := N / (2P²) with N = current sequence length. This is the recipe used by the local-global FLOP-matched schedule in Sec. 5.1 / Appendix C. Args: seq_len: S — longitud de la secuencia actual. window_size: W — tokens hacia atrás que cada query puede atender. device, dtype: del tensor Q para evitar copias de dispositivo. Returns: [1, 1, S, S] float — 0 en posiciones válidas, -inf en posiciones fuera de ventana (j < i − W). """ # i [S,1], j [1,S] → out-of-window where j < i - W i = torch.arange(seq_len, device=device).unsqueeze(1) # [S, 1] j = torch.arange(seq_len, device=device).unsqueeze(0) # [1, S] mask = torch.zeros(seq_len, seq_len, device=device, dtype=dtype) mask[j < (i - window_size)] = float("-inf") return mask.unsqueeze(0).unsqueeze(0) # [1, 1, S, S] def _apply_iha_pseudo_heads( self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """ IHA Steps 2-4 (Algorithm 1, Duvvuri et al. 2026, arXiv:2602.21371). Implements the paper-correct **seq-expand** approach: pseudo-heads are merged into the *sequence* dimension, yielding H heads each attending over S·P virtual tokens rather than the approximation of creating H·P independent heads over S tokens. Step 2 — Pseudo-head mixing across original heads (Alg. 1 lines 2-4): Q̃[h,p,s,:] = Σ_m α_Q[h,m,p] · Q[m,s,:] Step 3 — Interleave P pseudo-slots into sequence dimension (line 5): For each head h: sequence becomes (s=0,p=0), (s=0,p=1), …, (s=0,p=P-1), (s=1,p=0), …, (s=S-1,p=P-1) Virtual position of (s,p) is s·P + p. This is what enables up to P² distinct attention patterns per head (pseudo-query p₁ can attend to pseudo-key p₂ at the *same* original position s, which is impossible in the head-expand approximation). Compatibility with FlashAttention is preserved because Step 4 is standard scaled dot-product attention — no custom kernel needed. Requires flash_attn (imported lazily) for P>1; checked at forward call time, not here, so the function itself stays pure-PyTorch. P=1 Squeeze the unit P dimension — output shapes identical to input. Functionally equivalent to a learned linear recombination of heads (no seq expansion, no new parameters beyond alpha). P>1 (seq-expand) Q: [B, H_q, S, d] → [B, H_q, S·P, d] K/V: [B, H_kv, S, d] → [B, H_kv, S·P, d] GQA ratio H_q/H_kv is invariant (both expand in the seq dim). Args ---- q: [B, H_q, S, d_head] k: [B, H_kv, S, d_head] v: [B, H_kv, S, d_head] Returns ------- q_out: [B, H_q, S·P, d_head] (= [B, H_q, S, d] if P=1) k_out: [B, H_kv, S·P, d_head] (= [B, H_kv, S, d] if P=1) v_out: [B, H_kv, S·P, d_head] (= [B, H_kv, S, d] if P=1) """ P = self.iha_P B, H_q, S, d = q.shape H_kv = k.shape[1] # Step 2: learned linear combination across original heads. # alpha[h_out, h_in, p] × q[b, h_in, s, d] → [B, H_q, P, S, d] # einsum contracts h_in (index m) and distributes over pseudo-slot p. q_p = torch.einsum("hmp,bmsd->bhpsd", self.iha_alpha_q, q) # [B,H_q, P,S,d] k_p = torch.einsum("hmp,bmsd->bhpsd", self.iha_alpha_k, k) # [B,H_kv,P,S,d] v_p = torch.einsum("hmp,bmsd->bhpsd", self.iha_alpha_v, v) # [B,H_kv,P,S,d] if P == 1: # Squeeze unit P axis — shapes unchanged, all downstream paths intact. q_out = q_p.squeeze(2) # [B, H_q, S, d] k_out = k_p.squeeze(2) # [B, H_kv, S, d] v_out = v_p.squeeze(2) # [B, H_kv, S, d] else: # Step 3: interleave P pseudo-slots into the sequence dimension. # Layout of q_p: [B, H, P, S, d] # Desired: [B, H, S·P, d] with virtual pos s·P+p # # .transpose(2, 3) → [B, H, S, P, d] (swap P and S axes) # .reshape(...) → [B, H, S·P, d] (fuse S and P into one axis) # # The resulting index is: out[b, h, s·P+p, d] = q_p[b, h, p, s, d] # which is exactly the interleaving the paper describes. q_out = q_p.transpose(2, 3).reshape(B, H_q, S * P, d) k_out = k_p.transpose(2, 3).reshape(B, H_kv, S * P, d) v_out = v_p.transpose(2, 3).reshape(B, H_kv, S * P, d) return q_out, k_out, v_out def _apply_iha_collapse(self, attn_out: torch.Tensor) -> torch.Tensor: """ IHA Step 5 (Algorithm 1, line 7-8, Duvvuri et al. 2026). Collapses the P pseudo-slot axis back to one representation per real token position using the learned R matrix. **Paper formulation (seq-expand):** reshape(O, [H, N, P, d]) → (H, N, P, d) einsum('hp, hnpd → hnd', R, O) → (H, N, d) where R ∈ ℝ^{H×P} weights the contribution of each pseudo-slot p to the final head-h output. This is *not* a collapse over heads — it is a collapse over the P interleaved virtual tokens at each real position. Initialization (identity / Theorem 2 inclusion M ⊆ P_P): R[h, 0] = 1.0, R[h, p>0] = 0.0 for all h so at step 0 the output equals the p=0 pseudo-slot, which by the alpha initialization equals the original MHA head output. Args ---- attn_out : [B, S·P, H_q, d_head] Flash-attention output in HF layout (batch, seq, heads, dim). S·P is the expanded virtual sequence length. Returns ------- [B, S, H_q, d_head] One representation per real token, one per original head. """ B, SP, H_q, d = attn_out.shape P = self.iha_P S = SP // P # Step 7: separate the interleaved P pseudo-slots from the N real positions. # [B, S·P, H, d] → [B, S, P, H, d] out_struct = attn_out.reshape(B, S, P, H_q, d) # Step 8: weighted sum over pseudo-slots. # R: [H_q, P] einsum index p → 'hp,bsphd → bshd' return torch.einsum("hp,bsphd->bshd", self.iha_R, out_struct) # ── IHA seq-expand helpers ──────────────────────────────────────────────── def _build_iha_interleaved_rope( self, q: torch.Tensor, k: torch.Tensor, position_ids: torch.Tensor, inv_freq: torch.Tensor, attention_scaling: float, ) -> Tuple[torch.Tensor, torch.Tensor]: """ Build interleaved RoPE for the seq-expanded IHA sequence and apply it. The paper (§4, interleaving note) requires that virtual token (s, p) receives integer position s·P + p, giving each pseudo-slot a distinct rotary phase. This cannot be done with the precomputed cos/sin (which cover positions 0…S-1); we rebuild cos/sin inline from inv_freq exactly as _apply_repo_rope does for the REPO path. Interleaved position of virtual token (s, p): pos_iha[b, s·P + p] = position_ids[b, s] · P + p Construction: position_ids_iha[B, S·P]: • Expand position_ids[B, S] → [B, S, 1] · P + arange(P)[1,1,P] → [B, S, P], then reshape → [B, S·P] cos/sin computed inline (no autocast, float32 trig): freqs[B, S·P, rot_dim/2] = pos_iha[:, :, None] · inv_freq[None, None, :] emb = cat(freqs, freqs) → [B, S·P, rot_dim] cos, sin scaled by attention_scaling apply_rotary_pos_emb is then called with q[B, H_q, S·P, d] and k[B, H_kv, S·P, d] using the expanded cos/sin[B, S·P, rot_dim]. Args ---- q, k : already seq-expanded, [B, H, S·P, d] position_ids: original [B, S] integer positions (handles packing) inv_freq: rotary frequency vector [rot_dim/2] (from repo_rope_args) attention_scaling: float scale (from repo_rope_args) Returns ------- (q_rot, k_rot) with same shapes as input. """ # Don't read batch size B from q here — position_ids may have # pid_B=1 (broadcast-ready) while q has B=batch_size. # apply_rotary_pos_emb handles the broadcast internally. SP = q.shape[2] P = self.iha_P # ── Normalize position_ids to exactly 2-D [pid_B, S] ───────────────── # Trainers can pass position_ids in three shapes: # [S] 1-D: single sequence, no batch dim # [pid_B, S] 2-D: standard (may be pid_B=1 broadcast or pid_B=B) # [pid_B, S, k] 3-D: e.g. [1, S, 2] when trainers attach per-doc # metadata (doc_id, pos) or (cos_pos, sin_pos); # only the sequential index (dim -1 first col) is # needed for RoPE. # All three are normalized to 2-D so the arithmetic below is stable. pid = position_ids if pid.dim() == 1: pid = pid.unsqueeze(0) # [1, S] elif pid.dim() > 2: pid = pid[..., 0] # [pid_B, S] — take sequential dim # Build interleaved position indices [pid_B, S·P] # Virtual token (s, p) at index s·P+p receives position pid[b,s]·P + p. # p_offsets[P] = [0, 1, …, P-1] # pid.unsqueeze(-1) → [pid_B, S, 1] # * P + p_offsets.view(1, 1, P) → [pid_B, S, P] # reshape(-1, SP) → [pid_B, S·P] (-1 keeps pid_B intact even when 1) p_offsets = torch.arange(P, device=pid.device, dtype=pid.dtype) # [P] pos_iha = (pid.unsqueeze(-1) * P + p_offsets.view(1, 1, P)).reshape( -1, SP ) # [pid_B, S·P] # Compute cos/sin inline from inv_freq. # No torch.no_grad() wrapper: position_ids / inv_freq have no grad # (int64 and frozen buffer respectively), so no gradient path is # created regardless. Wrapping in no_grad() inside a compiled region # causes a Dynamo graph break (it cannot inline autograd context # managers mid-graph). inv_freq_f = inv_freq.to(dtype=torch.float32) freqs = pos_iha.float().unsqueeze(-1) * inv_freq_f # [pid_B, S·P, r/2] emb = torch.cat([freqs, freqs], dim=-1) # [pid_B, S·P, r] # Keep the RoPE scale on the same device as the FP32 trig inputs. A # Python float is otherwise lifted by AOTAutograd as a CPU f64 scalar; # under Inductor CUDAGraph replay that scalar can acquire a CUDA static # address while the generated backward still dereferences it from a CPU # C++ kernel. Materializing this single four-byte value on-device keeps # the graph device-consistent and the multiply remains fused. attention_scale_f = torch.as_tensor( attention_scaling, device=emb.device, dtype=emb.dtype, ) cos = (emb.cos() * attention_scale_f).to(q.dtype) # [pid_B, S·P, r] sin = (emb.sin() * attention_scale_f).to(q.dtype) # apply_rotary_pos_emb adds unsqueeze(1) internally for the head dim # and broadcasts over the batch dim, so pid_B=1 works for any B. q_rot, k_rot = apply_rotary_pos_emb(q, k, cos, sin) return q_rot, k_rot def _hola_rms_gamma( self, x: torch.Tensor, gamma: torch.Tensor, ) -> torch.Tensor: """ HOLA cache-read normalization used only for the memory subspace. Args: x: [B, T, H, d] tensor in flash-attn layout. gamma: [H, d] learned RMSNorm-γ scale. """ rms = x.float().pow(2).mean(dim=-1, keepdim=True).add(self.hola_score_eps).rsqrt() return (x.float() * rms * gamma.view(1, 1, gamma.shape[0], gamma.shape[1]).float()).to(x.dtype) def _hola_value_innovation_scores( self, v: torch.Tensor, gate_states: Optional[torch.Tensor], ) -> torch.Tensor: """ FA2-compatible HOLA-derived eviction proxy. Full HOLA writes by beta*||e|| where e is the committed residual of a recurrent delta state. This Transformer/FA2 path intentionally avoids a second attention pass and stock FA2 does not return the pre-dropout leave-one-out accumulators needed for the exact score. We therefore use a causal value-innovation proxy: score_t = || V_bundle(t) - V_bundle(t-1) ||_rms * gate_strength_t The score is computed per *real* token and aggregates all pseudo-slots, so TopK stores complete IHA episodes rather than isolated virtual slots. """ B, H_kv, SP, d_head = v.shape P = self.iha_P S_orig = SP // P v_tok = v.reshape(B, H_kv, S_orig, P, d_head).float() prev = torch.cat([torch.zeros_like(v_tok[:, :, :1]), v_tok[:, :, :-1]], dim=2) delta = v_tok - prev scores = delta.pow(2).mean(dim=(1, 3, 4)).add(self.hola_score_eps).sqrt() # [B, S] if gate_states is not None: gate = torch.sigmoid(gate_states).view(B, S_orig, self.config.num_attention_heads, d_head) gate_strength = gate.float().pow(2).mean(dim=(2, 3)).add(self.hola_score_eps).sqrt() scores = scores * gate_strength return scores def _hola_augmented_fa_call( self, q_cur: torch.Tensor, k_cur: torch.Tensor, v_cur: torch.Tensor, k_mem: Optional[torch.Tensor], v_mem: Optional[torch.Tensor], dropout: float, window_left: int, ) -> torch.Tensor: """ Single-FA2 HOLA read using augmented Q/K/V subspaces. Ordinary keys use the first d channels: / sqrt(D) = / sqrt(d) Memory keys use the second d channels plus one constant-bias channel: / sqrt(D) + b Values are padded with zeros and sliced back to d after FA2. The result is one softmax over [memory, current chunk], not a separate cache attention output. """ flash_attn_fn = self._iha_flash_attn_func B, Tq, H_q, d_head = q_cur.shape H_kv = k_cur.shape[2] base_aug = 2 * d_head + 1 d_aug = ((base_aug + 7) // 8) * 8 pad_tail = d_aug - base_aug scale = math.sqrt(float(d_aug) / float(d_head)) softmax_scale = 1.0 / math.sqrt(float(d_aug)) q_h = self._hola_rms_gamma(q_cur, self.hola_q_gamma) one_q = torch.ones(B, Tq, H_q, 1, device=q_cur.device, dtype=q_cur.dtype) q_parts = [q_cur * scale, q_h * scale, one_q] if pad_tail > 0: q_parts.append(torch.zeros(B, Tq, H_q, pad_tail, device=q_cur.device, dtype=q_cur.dtype)) q_aug = torch.cat(q_parts, dim=-1).contiguous() zeros_cur = torch.zeros_like(k_cur) zero_bias_cur = torch.zeros(B, k_cur.shape[1], H_kv, 1, device=k_cur.device, dtype=k_cur.dtype) k_cur_parts = [k_cur, zeros_cur, zero_bias_cur] if pad_tail > 0: k_cur_parts.append(torch.zeros(B, k_cur.shape[1], H_kv, pad_tail, device=k_cur.device, dtype=k_cur.dtype)) k_cur_aug = torch.cat(k_cur_parts, dim=-1) v_cur_pad = torch.zeros(B, v_cur.shape[1], H_kv, d_aug - d_head, device=v_cur.device, dtype=v_cur.dtype) v_cur_aug = torch.cat([v_cur, v_cur_pad], dim=-1) if k_mem is not None and k_mem.shape[1] > 0: k_mem_h = self._hola_rms_gamma(k_mem, self.hola_k_gamma) zeros_mem = torch.zeros_like(k_mem) bias = (self.hola_logit_bias.view(1, 1, H_kv, 1).to(k_mem.dtype) * math.sqrt(float(d_aug))) bias = bias.expand(B, k_mem.shape[1], H_kv, 1) k_mem_parts = [zeros_mem, k_mem_h, bias] if pad_tail > 0: k_mem_parts.append(torch.zeros(B, k_mem.shape[1], H_kv, pad_tail, device=k_mem.device, dtype=k_mem.dtype)) k_aug = torch.cat([torch.cat(k_mem_parts, dim=-1), k_cur_aug], dim=1).contiguous() v_mem_pad = torch.zeros(B, v_mem.shape[1], H_kv, d_aug - d_head, device=v_mem.device, dtype=v_mem.dtype) v_mem_aug = torch.cat([v_mem, v_mem_pad], dim=-1) v_aug = torch.cat([v_mem_aug, v_cur_aug], dim=1).contiguous() mem_len = k_mem.shape[1] else: k_aug = k_cur_aug.contiguous() v_aug = v_cur_aug.contiguous() mem_len = 0 # When local IHA is active, memory tokens are placed immediately before # the current chunk. Enlarging the left window by mem_len makes all # selected episodes visible while preserving causal masking inside the # current chunk. For global layers window_left is -1. if window_left >= 0: window = (mem_len + window_left, 0) else: window = (-1, -1) out_aug = flash_attn_fn( q_aug, k_aug, v_aug, dropout_p=dropout if self.training else 0.0, softmax_scale=softmax_scale, causal=True, window_size=window, ) return out_aug[..., :d_head].contiguous() def _iha_hola_seq_expand_flash_forward( self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, dropout: float, gate_states: Optional[torch.Tensor] = None, ) -> torch.Tensor: """ HOLA-derived episodic memory for the dense IHA seq-expand FA2 path. The sequence is processed in original-token chunks. At the beginning of a chunk, a fixed-size TopK memory is selected from strictly previous tokens using the value-innovation proxy. The chunk then performs a single FlashAttention-2 call over [memory KV, chunk KV]. This preserves causal semantics, keeps the memory frozen during the chunk, and avoids a second attention output. """ P = self.iha_P B, H_q, SP, d_head = q.shape H_kv = k.shape[1] S_orig = SP // P q_fa_all = q.transpose(1, 2).contiguous() # [B, S*P, H_q, d] k_fa_all = k.transpose(1, 2).contiguous() # [B, S*P, H_kv, d] v_fa_all = v.transpose(1, 2).contiguous() # [B, S*P, H_kv, d] scores = self._hola_value_innovation_scores(v, gate_states) # [B, S] chunk = max(1, int(self.hola_chunk_size)) mem_k = max(0, int(self.hola_memory_size)) window_left = self._resolve_iha_window() * P if self.iha_is_local else -1 outputs = [] for start in range(0, S_orig, chunk): end = min(start + chunk, S_orig) q_cur = q_fa_all[:, start * P : end * P] k_cur = k_fa_all[:, start * P : end * P] v_cur = v_fa_all[:, start * P : end * P] if mem_k > 0 and start > 0: k_select = min(mem_k, start) _, token_idx = torch.topk(scores[:, :start], k=k_select, dim=1) p_offsets = torch.arange(P, device=q.device, dtype=token_idx.dtype).view(1, 1, P) virtual_idx = (token_idx.unsqueeze(-1) * P + p_offsets).reshape(B, k_select * P) gather_idx = virtual_idx[:, :, None, None].expand(B, k_select * P, H_kv, d_head) k_mem = torch.gather(k_fa_all, dim=1, index=gather_idx).contiguous() v_mem = torch.gather(v_fa_all, dim=1, index=gather_idx).contiguous() else: k_mem = None v_mem = None outputs.append( self._hola_augmented_fa_call( q_cur=q_cur, k_cur=k_cur, v_cur=v_cur, k_mem=k_mem, v_mem=v_mem, dropout=dropout, window_left=window_left, ) ) return torch.cat(outputs, dim=1).contiguous() def _iha_seq_expand_flash_forward( self, q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, alpha: Optional[torch.Tensor], beta: Optional[torch.Tensor], use_affine: bool, dropout: float, position_ids: Optional[torch.Tensor], attention_mask: Optional[torch.Tensor] = None, gate_states: Optional[torch.Tensor] = None, ) -> Tuple[torch.Tensor, None]: """ Flash attention for the IHA seq-expand path (P>1). Bypasses the HF flash-attention wrapper entirely to avoid the three blockers identified in the architecture analysis: • attention_mask truncation in _upad_input (key silently clipped to S) • cu_seqlens / packing incompatibility (_prepare_from_posids shape mismatch) • attention_mask→additive-mask conversion that HF injects (unsupported for seq length S·P when the mask is shaped for S) Calls flash_attn_func on the dense fixed-shape path, or flash_attn_varlen_func only for true packed sequences without a padding mask. Right padding must not change IHA topology during fixed-length training; padding is handled by position ids and masked losses instead. Affine-Scaled Attention is also computed inline here so that Term 2 sums V over the same causal keys as flash attention on the S·P expanded sequence. Packing detection ----------------- Packing is assumed when batch_size == 1 AND position_ids is not a simple 0…S-1 range (i.e., contains document-boundary resets). This mirrors HF _is_packed_sequence logic. When packing is detected, cu_seqlens_iha are derived by multiplying per-document lengths by P. Sliding window -------------- For local IHA layers (iha_is_local=True), the window in the expanded sequence is W·P where W = S_orig // (2·P²) (paper Appendix C). Affine-Scaled Attention (inline) --------------------------------- When use_affine=True: output = α · flash_out + β · Σ_{j∈A(i)} V_j α, β are [B, H_q, S·P, 1] (already expanded over the seq dim before this call). The V sum is causal prefix-sum for global IHA, windowed for local IHA, and segmented when packed position_ids reset. Args ---- q, k, v : [B, H_q/kv, S·P, d] — after IHA expansion and RoPE alpha, beta: [B, H_q, S·P, 1] or None use_affine : whether Affine-Scaled Attention is active dropout : attention dropout probability S_orig : original (pre-expand) sequence length S position_ids: [B, S] or None — for packing detection + cu_seqlens gate_states : [B, S, H_q*d] or None — optional Gated-Attention gate logits used only by the HOLA-derived eviction proxy. Returns ------- (attn_out, None) attn_out: [B, S·P, H_q, d_head] (flash layout — batch, seq, heads, dim) """ # flash_attn functions are stored as instance attributes in __init__ # (self._iha_flash_attn_func / self._iha_flash_attn_varlen) so Dynamo # sees them as compile-time constants — no global lookup, no graph break. flash_attn_fn = self._iha_flash_attn_func flash_varlen_fn = self._iha_flash_attn_varlen P = self.iha_P B, H_q, SP, d_head = q.shape H_kv = k.shape[1] # flash_attn layout: [B, seq, heads, dim] q_fa = q.transpose(1, 2).contiguous() # [B, S·P, H_q, d] k_fa = k.transpose(1, 2).contiguous() # [B, S·P, H_kv, d] v_fa = v.transpose(1, 2).contiguous() # [B, S·P, H_kv, d] # ── Sliding window (local IHA layers) ──────────────────────────────── # Paper Appendix C: W := N / (2·P²) in original-sequence tokens. # In the expanded sequence each query can see W·P virtual keys to the # left. flash_attn window_size=(left, right) in virtual-token units. if self.iha_is_local: W_orig = self._resolve_iha_window() window = (W_orig * P, 0) else: window = (-1, -1) # ── Packing detection ───────────────────────────────────────────────── # Packing: batch_size==1 and position_ids contains document-boundary # resets (non-monotonic positions). When attention_mask is present we # keep the dense fixed-shape path: right padding should affect losses and # position ids, not IHA topology or compiled tensor shapes. _pos_for_pack = None if position_ids is not None and B == 1 and attention_mask is None: _p = position_ids[0] # [S] or [S, k] if _p.dim() > 1: _p = _p[:, 0] # [S] — take sequential index column _pos_for_pack = _p.long() # [S] # .all().item() would break the compiled graph; use bool() which Dynamo # can evaluate at trace time when the tensor is a compile-time constant, # and falls back to a graph-break-free path otherwise. is_packed = ( _pos_for_pack is not None and _pos_for_pack.numel() > 1 and not bool((torch.diff(_pos_for_pack.float()) >= 0).all()) ) # ── Packed / varlen path ────────────────────────────────────────────── if is_packed: # Derive cu_seqlens from document boundaries in the 1-D position # sequence. A new document starts wherever the position index resets # (pos[i] < pos[i-1]). pos0 = _pos_for_pack # [S] # Boundary mask: True at positions that start a new document. boundaries = torch.cat( [ torch.ones(1, device=pos0.device, dtype=torch.bool), pos0[1:] < pos0[:-1], # reset ] ) doc_starts = boundaries.nonzero(as_tuple=False).squeeze(1) # indices doc_lengths = torch.diff( torch.cat([doc_starts, torch.tensor([pos0.numel()], device=pos0.device)]) ) # [num_docs] # Scale each document length by P for the expanded sequence. doc_lengths_p = doc_lengths * P # [num_docs] cu_seqlens = torch.zeros( len(doc_lengths) + 1, device=q.device, dtype=torch.int32 ) cu_seqlens[1:] = doc_lengths_p.cumsum(0).to(torch.int32) max_seqlen_p = int(doc_lengths_p.max().item()) # Flatten batch dim (B=1 for packing): [1, S·P, H, d] -> [S·P, H, d] q_flat = q_fa.squeeze(0) k_flat = k_fa.squeeze(0) v_flat = v_fa.squeeze(0) out_flat = flash_varlen_fn( q_flat, k_flat, v_flat, cu_seqlens_q=cu_seqlens, cu_seqlens_k=cu_seqlens, max_seqlen_q=max_seqlen_p, max_seqlen_k=max_seqlen_p, dropout_p=dropout if self.training else 0.0, # FlashAttention's default is exactly head_dim**-0.5, which is # self.scaling. None avoids lifting a Python f64 scalar through # AOTAutograd while preserving identical attention logits. softmax_scale=None, causal=True, window_size=window, ) # [S·P, H_q, d] flash_out = out_flat.unsqueeze(0) # [1, S·P, H_q, d] # ── Dense (non-packed) path ─────────────────────────────────────────── else: if self.use_hola_memory and self.iha_P > 1 and not use_affine: # HOLA-derived memory is intentionally restricted to the dense # direct-FA2 seq-expand branch. Packed varlen and the affine # inline decomposition keep the original baseline path until a # dedicated fused kernel exposes the exact pre-dropout residual # statistics. flash_out = self._iha_hola_seq_expand_flash_forward( q, k, v, dropout=dropout, gate_states=gate_states, ) else: flash_out = flash_attn_fn( q_fa, k_fa, v_fa, dropout_p=dropout if self.training else 0.0, # Equivalent to self.scaling (head_dim**-0.5), without a # host scalar crossing the compiled CUDA graph boundary. softmax_scale=None, causal=True, window_size=window, ) # [B, S·P, H_q, d] # ── Affine-Scaled Attention inline (flash path) ─────────────────────── # Decomposition used here: # Term 1: α · flash_out (standard flash output) # Term 2: β · Σ_{j≤i} V_j (causal prefix value summary) # This intentionally restores the original local-IHA behavior: the # attention logits may use a sliding window, but the β branch supplies a # cheap global causal value context. For packed IHA, the prefix is # reset at each document boundary to avoid cross-document leakage. # α, β arrive pre-expanded to [B, H_q, S·P, 1]; permute to [B,S·P,H,1] # to match the flash output layout [B, S·P, H_q, d]. if use_affine: # Compute the β value-sum in KV-head space first. For GQA this is # exact because the valid-key sum is linear and independent of the # repeated query-head view: # Sum(repeat_kv(V)) == repeat_kv(Sum(V)). # The broadcast to H_q happens only in the final affine combine. # β uses the full causal prefix, not the local flash window. # Keep the sum in KV-head space first; GQA expansion happens only # in the final affine combine below. if is_packed: value_sum_kv = _segmented_causal_value_sum( v, segment_lengths=doc_lengths_p, window_left=None, ) # [1,H_kv,S·P,d] else: value_sum_kv = _causal_value_sum( v, query_len=SP, window_left=None, ) # [B,H_kv,S·P,d] alpha_t = alpha.permute(0, 2, 1, 3) # [B,S·P,H_q,1] beta_t = beta.permute(0, 2, 1, 3) # [B,S·P,H_q,1] value_sum_t = value_sum_kv.transpose(1, 2) # [B,S·P,H_kv,d] if self.num_key_value_groups == 1 or H_q == H_kv: attn_out = alpha_t * flash_out + beta_t * value_sum_t else: flash_g = flash_out.reshape( B, SP, H_kv, self.num_key_value_groups, d_head ) alpha_g = alpha_t.reshape(B, SP, H_kv, self.num_key_value_groups, 1) beta_g = beta_t.reshape(B, SP, H_kv, self.num_key_value_groups, 1) attn_out = ( alpha_g * flash_g + beta_g * value_sum_t.unsqueeze(3) ).reshape(B, SP, H_q, d_head) attn_out = torch.nn.functional.dropout( attn_out, p=dropout, training=self.training ) else: attn_out = flash_out return attn_out, None def _apply_momentum_attention( self, q: torch.Tensor, k: torch.Tensor, ) -> Tuple[torch.Tensor, torch.Tensor]: if not self.use_momentum_attention or self.momentum_gamma == 0.0: return q, k dq = causal_first_difference(q) dk = causal_first_difference(k) gamma = torch.as_tensor( self.momentum_gamma, device=q.device, dtype=torch.float32, ) q_new = q + gamma * dq k_new = k + gamma * dk return q_new, k_new def _apply_mea_head_mixing( self, k: torch.Tensor, v: torch.Tensor, _force_standard: bool = False, ) -> Tuple[torch.Tensor, torch.Tensor]: if not self.use_mea_attention: return k, v if self.use_iha and self.iha_P > 1 and not _force_standard: # Preserve the IHA pseudo-slot factorization: MEA mixes the KV # component heads inside each pseudo-slot, but never mixes # information across different pseudo indices. # NOTE: this branch is reached only by the head-expand path (P>1 # but _force_standard=False). The seq-expand path (P>1) calls # with _force_standard=True and falls through to the standard branch. k_mixed = head_linear_compose_pseudo( k, self.mea_key_mix, self.iha_P ).contiguous() v_mixed = head_linear_compose_pseudo( v, self.mea_value_mix, self.iha_P ).contiguous() else: # Standard path — also used by IHA seq-expand (P>1, _force_standard=True). # K/V may be [B, H_kv, S·P, d]; head_linear_compose operates on the # head dimension and leaves the sequence dimension untouched. k_mixed = head_linear_compose(k, self.mea_key_mix).contiguous() v_mixed = head_linear_compose(v, self.mea_value_mix).contiguous() return k_mixed, v_mixed def _apply_lucid_preconditioner( self, k: torch.Tensor, v: torch.Tensor, attention_mask: Optional[torch.Tensor], ) -> torch.Tensor: if not self.use_lucid_attention: return v.contiguous() key_rn = rms_key_unit_norm(k, eps=self.lucid_attention_eps) logits = ( torch.matmul(key_rn, key_rn.transpose(-1, -2)) * self.scaling - self.sqrt_head_dim ) prec = torch.tril(torch.exp(logits)) kv = infer_key_validity(attention_mask, k.shape[-2], k.shape[1]) if kv is not None: prec = prec * (kv.unsqueeze(-1) & kv.unsqueeze(-2)).to(prec.dtype) eye = torch.eye(prec.shape[-1], device=prec.device, dtype=prec.dtype).view( 1, 1, prec.shape[-1], prec.shape[-1] ) prec = prec + eye * (1.0 - prec.diagonal(dim1=-2, dim2=-1).unsqueeze(-1)) result = ( torch.linalg.solve_triangular( prec, v.float(), upper=False, unitriangular=True ) .to(v.dtype) .contiguous() ) return result def _apply_directional_routing( self, attn_out: torch.Tensor, hidden_states: torch.Tensor, ) -> torch.Tensor: """ Directional suppression at position C (post-XSA, pre-reshape). Args: attn_out: [B, S, H, d_head] — output after XSA and RMSNorm. hidden_states: [B, S, hidden_size] — pre-FAN residual stream, used as router input (same as paper's x_i). Returns: [B, S, H, d_head] with selected directional components suppressed. """ # ── Router: one routing decision per sequence ───────────────────── # Mean-pool over sequence dimension S → [B, hidden_size]. # The router is sequence-level (not token-level), matching the paper. # This produces one suppression pattern per input sequence, # which the paper shows learns domain-adaptive behavior in early # layers and fixed syntactic pruning in late layers. pooled = hidden_states.mean(dim=1) # [B, hidden_size] logits = self.direction_router(pooled) # [B, H*K] r = torch.sigmoid(self.directional_routing_temp * logits) r = r.view( hidden_states.shape[0], self.config.num_attention_heads, self.directional_routing_k, ) # [B, H, K] # Expand over sequence for broadcasting: [B, 1, H, K] r = r.unsqueeze(1) # ── Unit-normalize direction vectors ────────────────────────────── # Normalize at forward time, not at init, following the paper. # d: [H, K, d_head] → unit norm along d_head dimension. d = F.normalize(self.direction_vecs, dim=-1) # [H, K, d_head] # ── Directional suppression ─────────────────────────────────────── # attn_out: [B, S, H, d_head] # For each head h and direction k: # proj_{h,k} = (o_h · d_{h,k}) scalar per (B, S, H, K) # suppress = r_{h,k} · proj_{h,k} · d_{h,k} # o'_h = o_h - Σ_k suppress_{h,k} # # proj: einsum over d_head dimension # attn_out [B, S, H, D] × d [H, K, D] → [B, S, H, K] proj = torch.einsum("bshd,hkd->bshk", attn_out, d) # [B, S, H, K] # r [B, 1, H, K] × proj [B, S, H, K] → [B, S, H, K] weighted = r * proj # [B, S, H, K] # Σ_k weighted_{h,k} · d_{h,k}: # weighted [B, S, H, K] × d [H, K, D] → [B, S, H, D] suppression = torch.einsum("bshk,hkd->bshd", weighted, d) result = attn_out - suppression return result def forward( self, hidden_states: torch.Tensor, position_embeddings: tuple[torch.Tensor, torch.Tensor], attention_mask: Optional[torch.Tensor] = None, first_layer_fan: Optional[torch.Tensor] = None, repo_rope_args: Optional[Tuple[torch.Tensor, float]] = None, position_ids: Optional[torch.LongTensor] = None, key_value_states: Optional[Tuple[torch.Tensor, torch.Tensor]] = None, **kwargs: Unpack[FlashAttentionKwargs], ) -> tuple[torch.Tensor, Optional[torch.Tensor], Optional[torch.Tensor]]: input_shape = hidden_states.shape[:-1] if key_value_states is None: key_states = hidden_states value_states = hidden_states else: key_states, value_states = key_value_states h_fan_q = self.fan_layer(hidden_states) if key_value_states is None: h_fan_k = h_fan_q h_fan_v = h_fan_q else: h_fan_k = self.fan_layer(key_states) h_fan_v = self.fan_layer(value_states) if self.use_fan_residual and first_layer_fan is not None: h_fan_q = self.lambda_1 * first_layer_fan + self.lambda_2 * h_fan_q h_fan_k = self.lambda_1 * first_layer_fan + self.lambda_2 * h_fan_k h_fan_v = self.lambda_1 * first_layer_fan + self.lambda_2 * h_fan_v current_layer_fan = h_fan_q.clone() query_shape = (*input_shape, self.config.num_attention_heads, self.head_dim) kv_shape = (*input_shape, self.num_mea_component_heads, self.head_dim) q_raw, gate = torch.chunk( self.q_proj(h_fan_q).view( *input_shape, self.config.num_attention_heads, self.head_dim * 2 ), 2, dim=-1, ) gate = gate.reshape(*input_shape, -1) q = self.q_norm(q_raw.view(query_shape)).transpose(1, 2) k = self.k_norm(self.k_proj(h_fan_k).view(kv_shape)).transpose(1, 2) v = self.v_proj(h_fan_v).view(kv_shape).transpose(1, 2) # ── IHA: cross-head pseudo-head mixing (Duvvuri et al., 2026) ──────── # Position: post active Q/K norm (RMSNorm or RMSNorm), pre-RoPE/REPO. # Applies learned linear combination of all H heads into pseudo-heads # before positional encoding (paper Alg. 1, Steps 2-4). # # P=1 → shapes unchanged; all downstream paths fully intact. # P>1 → seq-expand (paper-correct) only on layers where self.use_iha=True: # Q : [B, H_q, S, d] → [B, H_q, S·P, d] # K/V : [B, H_kv, S, d] → [B, H_kv, S·P, d] # Virtual token (s, p) is at index s·P+p in the expanded dim. # GQA ratio H_q / H_kv is invariant (both expand in seq, not heads). # For pattern token 'G' with iha_global_layers_use_iha=False, # self.use_iha=False in this layer and the standard global path runs. _iha_seq_expand = self.use_iha and self.iha_P > 1 _S_orig = q.shape[2] # original sequence length before any IHA expansion if self.use_iha: q, k, v = self._apply_iha_pseudo_heads(q, k, v) # ── RoPE / REPO ─────────────────────────────────────────────────────── cos, sin = position_embeddings _repo_grape_kernel_used = False if self.use_repo: # REPO path: f_ϕ predicts continuous per-head positions from the # residual stream, then cos/sin are built inline from those positions # so the rotation is differentiable w.r.t. REPOModule parameters. # inv_freq and attention_scaling arrive via repo_rope_args, sourced # directly from rotary_emb at forward time — no buffer on this module, # no meta-tensor issue on lm_eval / to(device) paths. # (Li et al., 2026, §3.2 — Eq. 6–7) z = self.repo_module(hidden_states) # [B, H, S], raw REPO assignment inv_freq, attn_scaling = repo_rope_args # The optional kernel consumes raw REPO assignments and performs # the linear-blend coordinate, IHA P=2 sequence expansion, GRAPE # phases, GQA circular projection, and optional momentum in one # autograd-aware Triton boundary. Q/K are intentionally normalized # before this call: IHA mixes normalized heads, so moving RMSNorm # into the kernel after that mixing would change the model. _repo_grape_sequence_factor = self.iha_P if _iha_seq_expand else 1 if ( self.use_repo_grape and self.use_repo_grape_triton and _cce_repo_grape is not None and _cce_repo_grape_supported is not None and _repo_grape_sequence_factor in (1, 2) and _cce_repo_grape_supported( q, k, z, inv_freq, self.repo_grape.position_blend_alpha, self.repo_grape.freq_log_scale, sequence_pseudo_factor=_repo_grape_sequence_factor, ) ): _kernel_position_ids = position_ids if ( _kernel_position_ids is not None and _kernel_position_ids.dim() > 2 ): # Match transform_positions and the IHA helpers: the first # metadata column is the runtime sequential position. _kernel_position_ids = _kernel_position_ids[..., 0] _kernel_gamma = ( self.momentum_gamma if self.use_momentum_attention else 0.0 ) q, k = _cce_repo_grape( q, k, z, _kernel_position_ids, inv_freq, self.repo_grape.position_blend_alpha, self.repo_grape.freq_log_scale, attn_scaling, sequence_pseudo_factor=_repo_grape_sequence_factor, momentum_gamma=_kernel_gamma, output_dtype=q.dtype, ) _repo_grape_kernel_used = True # REPO-GRAPE consumes u, not raw z. There is intentionally no mode # flag: when REPO-GRAPE is active, transform_positions always applies # the learned linear blend # # u_i^(h) = i + alpha_h * (z_i^(h) - i). # # Pure REPO (without GRAPE) keeps the original paper behavior and # therefore continues to use z directly. if not _repo_grape_kernel_used: if self.use_repo_grape: # Keep the exact canonical runtime positions alongside u. # They are needed only by the weak circular prior on shared # GQA keys and are never materialized once per query head. repo_coords, repo_reference_positions = ( self.repo_grape.transform_positions( z, position_ids, return_reference=True ) ) else: repo_coords = z repo_reference_positions = None if not _repo_grape_kernel_used and _iha_seq_expand: # IHA is applied *after* the coordinate map. Thus REPO-GRAPE # uses u_{i,p}=P*u_i+p, while pure REPO keeps its previous # z_{i,p}=P*z_i+p behavior. ``repo_coords`` is transient and this # expanded tensor replaces, rather than accompanies, the old # z_expanded allocation. P = self.iha_P p_offsets = torch.arange( P, device=repo_coords.device, dtype=repo_coords.dtype ) coords_expanded = ( repo_coords.unsqueeze(-1) * P + p_offsets.view(1, 1, 1, P) ).reshape(repo_coords.shape[0], repo_coords.shape[1], _S_orig * P) if self.use_repo_grape: # Match the canonical prior to the same IHA pseudo-slot: # u_{i,p}=P*u_i+p => i_{i,p}=P*i+p. reference_expanded = ( repo_reference_positions.unsqueeze(-1) * P + p_offsets.view(1, 1, P) ).reshape(repo_reference_positions.shape[0], _S_orig * P) q, k = self.repo_grape.apply_multiplicative( q, k, coords_expanded, inv_freq, attn_scaling, reference_positions=reference_expanded, ) else: q, k = _apply_repo_rope( q, k, coords_expanded, inv_freq, attn_scaling ) elif not _repo_grape_kernel_used: # P=1/no-IHA: ``repo_coords`` is u for REPO-GRAPE and z for # the original REPO-RoPE path. if self.use_repo_grape: q, k = self.repo_grape.apply_multiplicative( q, k, repo_coords, inv_freq, attn_scaling, reference_positions=repo_reference_positions, ) else: q, k = _apply_repo_rope( q, k, repo_coords, inv_freq, attn_scaling ) elif _iha_seq_expand: # IHA seq-expand + standard integer RoPE: # Build interleaved positions pos_iha[s·P+p] = position_ids[s]·P + p # and compute cos/sin inline so that each pseudo-slot receives a # distinct rotary phase — paper §4 interleaving note. # inv_freq is available via repo_rope_args (NeoLLMModel always passes # it when use_iha=True and P>1, regardless of use_repo). inv_freq, attn_scaling = repo_rope_args _pid = position_ids if _pid is None: # Fall back to 0…S-1 when position_ids were not threaded through. _pid = ( torch.arange(_S_orig, device=q.device) .unsqueeze(0) .expand(q.shape[0], -1) ) q, k = self._build_iha_interleaved_rope(q, k, _pid, inv_freq, attn_scaling) else: # Standard path: integer positions pre-computed by NeoLLMModel. q, k = apply_rotary_pos_emb(q, k, cos, sin) if not _repo_grape_kernel_used: q, k = self._apply_momentum_attention(q, k) # ── MEA head mixing ─────────────────────────────────────────────────── # For IHA seq-expand (P>1): K/V are [B, H_kv, S·P, d] — same number of # heads as without IHA, just a longer sequence. Use the standard # head_linear_compose path; the pseudo-slot-aware path (head_linear_compose_pseudo) # was designed for the head-expand layout and would mismatch here. # For P=1 or no IHA: the existing branch logic in _apply_mea_head_mixing # is unchanged and handles both the pseudo-slot-aware and normal cases. # Standard FlashAttention layers use the dense right-padded training # contract even inside a mixed IHA schedule. Keep ``attention_mask`` # itself untouched for IHA packing detection and auxiliary operators; # only the backend argument is normalized to ``None``. backend_attention_mask = attention_mask if ( self.training and self.config._attn_implementation in {"flash_attention_2", "flash_attention_3"} ): backend_attention_mask = None if _iha_seq_expand: k, v = self._apply_mea_head_mixing(k, v, _force_standard=True) else: k, v = self._apply_mea_head_mixing(k, v) v = self._apply_lucid_preconditioner(k, v, attention_mask) # Capture v_ref for XSA after MEA mixing and LUCID preconditioning. # This is the vector that actually participated in SDPA aggregation. v_ref = v if self.use_xsa else None # ── REPO-GOAT factorised log-prior ─────────────────────────────────── # Appends Q/K prior channels and zero V channels immediately before the # attention backend. The output is sliced back to head_dim after SDPA, # so all downstream modules remain shape-compatible. repo_goat_prior_dim = 0 if self.use_repo_goat_prior: q, k, v, repo_goat_prior_dim = self.repo_goat_prior.append_prior_subspace( q, k, v, position_ids=position_ids, ) # Normalize FlashAttention inputs for every topology, not only the # non-IHA fallback. RMSNorm produces FP32 under BF16 autocast; IHA's # pseudo-head einsums happened to cast it back implicitly, whereas the # standard route reached the Transformers wrapper in FP32. An explicit # dispatch-boundary cast gives IHA and non-IHA identical dtype behavior, # avoids hidden backend casts/warnings, and keeps FP32 work where it is # intentional (normalization, RoPE/REPO and momentum calculations). if self.config._attn_implementation in { "flash_attention_2", "flash_attention_3", }: flash_dtype = _flash_target_dtype(q, self) if flash_dtype is not None: q = q.to(flash_dtype) k = k.to(flash_dtype) v = v.to(flash_dtype) # ── IHA: local sliding-window mask for non-seq-expand path (P=1) ───── # For P>1 (seq-expand): the sliding window is expressed as window_size # in flash_attn_func and handled inside _iha_seq_expand_flash_forward. # This block only applies when P=1 (or use_iha=False) and the standard # HF attention backend is in use. # # Para capas marcadas como 'L' en el pattern (iha_is_local=True): # suma la máscara de ventana deslizante W a la máscara causal existente. # Esto restringe cada query a solo los últimos W tokens, reduciendo el # costo de atención de O(S²) a O(S·W) por head y FLOP-matcheando contra # global attention estándar con el schedule 4L+1G del paper. # # Compatibilidad FlashAttention-2 (paper §2 y §4, Algoritmo 1): # IHA preserva el operador de atención estándar — FlashAttention lo # recibe sin modificaciones. La máscara aditiva -inf es convertida # internamente por el wrapper HF a índices de bloque para flash2. # Para sdpa/eager la suma directa de -inf funciona nativamente. # Con P=1 iha_is_local=False siempre (ningún overhead). # En capas 'G' con iha_global_layers_use_iha=False, self.use_iha=False y # esta rama tampoco se activa: la capa es global estándar, paper-faithful. # flash_attention_2/3 keeps the 2D padding mask contract and receives # the locality constraint through `sliding_window`; eager/sdpa still # need the explicit additive band mask. flash_local_sliding_window = None if self.use_iha and self.iha_is_local and not _iha_seq_expand: _S = q.shape[2] # S real (puede diferir de max_pos con packing) _W = self._resolve_iha_window() if self.config._attn_implementation in { "flash_attention_2", "flash_attention_3", }: flash_local_sliding_window = _W else: _band = self._build_iha_local_mask(_S, _W, q.device, q.dtype) attention_mask = ( _band if attention_mask is None else attention_mask + _band ) # ── Affine-Scaled Attention ─────────────────────────────────────── # Active whenever use_affine_scaled_attention=True, regardless of # attention backend. Two code paths — same math, different execution: # eager : full weight access, attn_weights_pre/post_affine captured. # flash/sdpa: α·backend_out + β·Σ valid V, no weight tensors materialised. alpha = None beta = None use_affine = self.use_affine_scaled_attention if use_affine: alpha = linear_clipping(self.alpha_proj(hidden_states)) # [B, S, H] alpha = alpha.permute(0, 2, 1).unsqueeze(-1) # [B, H, S, 1] N = k.shape[-2] beta = (self.alpha_ma.to(alpha.dtype) - alpha) / max(N, 1) if self.training: with torch.no_grad(): batch_mean = alpha.mean(dim=(0, 2), keepdim=True) self.alpha_ma.copy_( self.affine_momentum * self.alpha_ma + (1.0 - self.affine_momentum) * batch_mean ) # ── IHA: expand alpha/beta to match the attention sequence ──────── # P=1 or no IHA: no expansion needed (shapes already correct). # # seq-expand (P>1): SDPA operates over S·P virtual tokens. # alpha_proj produces one scalar per *original* token per head # (computed from hidden_states before IHA expansion, so shape is # [B, H_q, S, 1]). Each original token's alpha/beta applies # uniformly to all P pseudo-slots — alpha is a property of the # real token, not of the virtual slot. # Expansion: repeat_interleave over the sequence dim (dim=2), NOT # the head dim (dim=1) which was the head-expand approximation. # EMA stats remain on [B, H, S, 1] (pre-expansion) — no change. if _iha_seq_expand: alpha = alpha.repeat_interleave(self.iha_P, dim=2) # [B,H_q,S·P,1] beta = beta.repeat_interleave(self.iha_P, dim=2) # [B,H_q,S·P,1] # Recompute beta with N = S·P for the expanded attention axis. # alpha_ma shape [1,H,1,1] broadcasts correctly against expanded alpha. N_expanded = k.shape[-2] # S·P after IHA expansion beta = (self.alpha_ma.to(alpha.dtype) - alpha) / max(N_expanded, 1) if _iha_seq_expand: # ── IHA seq-expand: bypass HF wrapper, call flash_attn directly ── # The HF flash-attention wrapper (_flash_attention_forward via # ALL_ATTENTION_FUNCTIONS) has three hard blockers for the S·P expanded # sequence: # 1. _upad_input silently truncates K/V to mask length S (line 45-46) # 2. _prepare_from_posids shape-mismatches packed sequences # 3. attention_mask shaped [B,S] is incompatible with S·P tokens # _iha_seq_expand_flash_forward calls flash_attn_func / # flash_attn_varlen_func directly, handling packing, sliding window, # and the Affine-Scaled inline decomposition (α·flash + β·Σ valid V). # alpha/beta are [B, H_q, S·P, 1] at this point (already expanded). attn_out, attn_weights = self._iha_seq_expand_flash_forward( q, k, v, alpha=alpha if use_affine else None, beta=beta if use_affine else None, use_affine=use_affine, dropout=0.0 if not self.training else self.attention_dropout, position_ids=position_ids, attention_mask=attention_mask, gate_states=gate, ) elif use_affine: if self.config._attn_implementation == "eager": # Eager: materialises softmax weights for the affine-scaled path. attn_out, attn_weights = affine_scaled_eager_attention_forward( self, q, k, v, attention_mask, scaling=self.scaling, alpha=alpha, beta=beta, dropout=0.0 if not self.training else self.attention_dropout, **kwargs, ) else: # Flash / SDPA: valid-key value sum, no weight tensors. backend_kwargs = kwargs if flash_local_sliding_window is not None: backend_kwargs = { **kwargs, "sliding_window": flash_local_sliding_window, } attn_out, attn_weights = affine_scaled_flash_attention_forward( self, q, k, v, backend_attention_mask, scaling=self.scaling, alpha=alpha, beta=beta, dropout=0.0 if not self.training else self.attention_dropout, **backend_kwargs, ) else: if self.config._attn_implementation == "eager": attn_fn = eager_attention_forward elif self.config._attn_implementation in { "flash_attention_2", "flash_attention_3", }: attn_fn = neollm_flash_attention_forward else: attn_fn = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation] backend_kwargs = kwargs if flash_local_sliding_window is not None: backend_kwargs = { **kwargs, "sliding_window": flash_local_sliding_window, } attn_out, attn_weights = attn_fn( self, q, k, v, backend_attention_mask, dropout=0.0 if not self.training else self.attention_dropout, # SDPA's None default is exactly head_dim**-0.5. Avoid lifting # that Python f64 scalar into the compiled backward graph. scaling=( None if self.config._attn_implementation == "sdpa" else self.scaling ), **backend_kwargs, ) if repo_goat_prior_dim: # Remove the zero-valued prior value channels. The prior has already # affected the attention probabilities through Q/K logits; these extra # output channels carry no value information by construction. attn_out = attn_out[..., : self.head_dim].contiguous() # ── IHA Step 5: collapse P pseudo-slot axis → one output per real token ─ # seq-expand (P>1): _iha_seq_expand_flash_forward returns [B, S·P, H, d]. # _apply_iha_collapse reshapes to [B, S, P, H, d] and applies # R[H, P] via einsum 'hp,bsphd→bshd' → [B, S, H, d]. # This is the paper Alg. 1 Step 7-8 (R ∈ ℝ^{H×P}, collapses P slots). # P=1: not executed — shapes already correct (R would be identity anyway). if self.use_iha and self.iha_P > 1: attn_out = self._apply_iha_collapse(attn_out) attn_out = attn_out.reshape(*input_shape, -1, self.head_dim) if self.use_mea_attention: attn_out = self.mea_output_norm(attn_out) # ── Exclusive Self Attention (position B, pre-routing) ──────────────── # Removes auto-position component before directional routing so that # direction_vecs specialize exclusively in cross-domain interference. # # v_ref is captured post-MEA, post-LUCID — the value representation that # actually participated in SDPA. For seq-expand IHA (P>1) v_ref has # shape [B, H_kv, S·P, d]; after repeat_kv and transpose it becomes # [B, S·P, H_q, d]. We apply the same R collapse as for attn_out so # that both tensors live in the same [B, S, H_q, d] space before the # orthogonal projection. For P=1 or no IHA the block is unchanged. if self.use_xsa and v_ref is not None: v_ref_expanded = repeat_kv( v_ref, self.num_key_value_groups ) # [B, H_q, S(·P), d] v_ref_t = v_ref_expanded.transpose(1, 2) # [B, S(·P), H_q, d] if _iha_seq_expand: # seq-expand: collapse [B, S·P, H_q, d] → [B, S, H_q, d] using # the same R[H_q, P] as attn_out, preserving the shared projection # subspace needed for the XSA dot product to be meaningful. v_ref_t = self._apply_iha_collapse(v_ref_t) elif self.use_iha and self.iha_P > 1: # head-expand legacy path (P>1 but seq-expand inactive — kept for # backward compatibility if the architecture is ever reconfigured). v_ref_t = self._apply_iha_collapse(v_ref_t) v_ref_t = v_ref_t.to(attn_out.dtype) proj = (attn_out * v_ref_t).sum(dim=-1, keepdim=True) norm_sq = torch.ops.aten.clamp_min.default( (v_ref_t * v_ref_t).sum(dim=-1, keepdim=True), self.xsa_eps, ) xsa_comp = (proj / norm_sq) * v_ref_t attn_out = attn_out - xsa_comp # ── Directional Routing (position C, post-XSA, pre-reshape) ────── # Suppresses cross-domain interference directions from the head output. # Operates on [B, S, H, d_head] before reshape and o_proj. # When use_xsa=False: directions span full head-space (no XSA pre-clean). # When use_directional_routing=False: this block is skipped entirely. if self.use_directional_routing: attn_out = self._apply_directional_routing(attn_out, hidden_states) # ── Reshape → o_proj → Gated Attention gate → dropout ──────────── attn_out_flat = attn_out.reshape(*input_shape, -1).contiguous() attn_out_gated = self.o_proj(attn_out_flat * torch.sigmoid(gate)) attn_out_gated = self.dropout(attn_out_gated) return attn_out_gated, attn_weights, current_layer_fan class PolyNorm(nn.Module): def __init__( self, eps: float = 1e-6, proj_eps: float = 1e-6, exclusive_init: float = 0.5, exclusive: bool = True, ): super().__init__() self.weight = nn.Parameter(torch.ones(3) / 3) self.bias = nn.Parameter(torch.zeros(1)) self.eps = eps self.exclusive = exclusive if exclusive: self.proj_eps = proj_eps # Dos fuerzas exclusivas aprendibles en (0, 1), una por rama de orden alto. # Se parametrizan con logits para que sigmoid mantenga alpha ∈ (0, 1). exclusive_init = float(min(max(exclusive_init, 1e-4), 1.0 - 1e-4)) init = torch.full((2,), exclusive_init, dtype=torch.float32) self.exclusive_logits = nn.Parameter(torch.logit(init)) def _norm(self, x): return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) def _exclusive(self, branch, ref, alpha, x1_f, ref_norm_sq): """ Elimina de `branch` la componente alineada con `ref` (la rama lineal x1), ponderada por alpha ∈ (0, 1), y renormaliza el resultado. branch := _norm(branch - alpha · proj_{ref}(branch)) El denominador ref_norm_sq se pasa precalculado para evitar duplicarlo cuando se llama dos veces por forward (una vez para x2, otra para x3). """ branch_f = branch.float() dot = (branch_f * x1_f).sum(dim=-1, keepdim=True) proj_coeff = (dot / ref_norm_sq).to(branch.dtype) out = branch - alpha.to(branch.dtype) * proj_coeff * ref return self._norm(out) def forward( self, x: torch.Tensor, dropout_p: float = 0.0, ) -> torch.Tensor: if _cce_polynorm is not None: if ( self.training and torch.is_grad_enabled() and torch.compiler.is_compiling() and _cce_polynorm_uses_cute is not None and _cce_polynorm_uses_cute( x, self.weight, self.bias, eps=float(self.eps), exclusive_logits=getattr(self, "exclusive_logits", None), dropout_p=dropout_p, ) ): return _activation_checkpoint( _cce_polynorm, x, self.weight, self.bias, use_reentrant=False, eps=float(self.eps), proj_eps=float(getattr(self, "proj_eps", 1e-6)), exclusive_logits=getattr(self, "exclusive_logits", None), dropout_p=dropout_p, ) return _cce_polynorm( x, self.weight, self.bias, eps=float(self.eps), proj_eps=float(getattr(self, "proj_eps", 1e-6)), exclusive_logits=getattr(self, "exclusive_logits", None), dropout_p=dropout_p, ) # Caché de potencias: x_sq reutilizado en x1 y x2; x_cu = x·x_sq evita pow(3) x_sq = x.pow(2) x_cu = x * x_sq # Tres ramas normalizadas x1 = x * x_sq.mean(-1, keepdim=True).add(self.eps).rsqrt() x2 = x_sq * (x_sq * x_sq).mean(-1, keepdim=True).add(self.eps).rsqrt() x3 = x_cu * (x_cu * x_cu).mean(-1, keepdim=True).add(self.eps).rsqrt() if self.exclusive: # Fuerzas exclusivas aprendibles alpha2, alpha3 = torch.sigmoid(self.exclusive_logits).unbind() # Precalcular ref (x1) en fp32 y su norma al cuadrado — compartido por x2 y x3 x1_f = x1.float() ref_norm_sq = x1_f.pow(2).sum(-1, keepdim=True).clamp_min(self.proj_eps) # Ortogonalización parcial de las ramas de orden alto respecto a la lineal x2 = self._exclusive(x2, x1, alpha2, x1_f, ref_norm_sq) x3 = self._exclusive(x3, x1, alpha3, x1_f, ref_norm_sq) output = ( self.weight[0] * x3 + self.weight[1] * x2 + self.weight[2] * x1 + self.bias ) if dropout_p: output = F.dropout(output, p=dropout_p, training=True) return output class NeoLLMMLP(nn.Module): """MLP with FANformer integration and Learnable Multipliers.""" def __init__(self, config): super().__init__() self.fan_layer = FANLayer( hidden_size=config.hidden_size, fan_ratio=getattr(config, "fan_ratio_ffn", 0.0625), ) fan_dim = config.hidden_size + int( config.hidden_size * getattr(config, "fan_ratio_ffn", 0.0625) ) self.gate_proj = LinearWithMultipliers( fan_dim, config.intermediate_size, bias=False, use_row_multiplier=True, use_column_multiplier=False, enable_multipliers=getattr(config, "use_learnable_multipliers", True), ) self.up_proj = nn.Linear(fan_dim, config.intermediate_size, bias=False) self.down_proj = LinearWithMultipliers( config.intermediate_size, config.hidden_size, bias=False, use_row_multiplier=True, use_column_multiplier=True, enable_multipliers=getattr(config, "use_learnable_multipliers", True), ) self.act_fn = PolyNorm( exclusive_init=0.00, exclusive=getattr(config, "polynorm_exclusive", True) ) self.dropout = nn.Dropout(config.dropout_rate) def forward( self, x: torch.Tensor, ) -> torch.Tensor: x_fan = self.fan_layer(x) gate_out = self.gate_proj(x_fan) up_out = self.up_proj(x_fan) dropout_p = float(self.dropout.p) if self.training else 0.0 act_out = self.act_fn(gate_out, dropout_p=dropout_p) act_x_up = act_out * up_out result = self.down_proj(act_x_up) return result class NeoLLMDecoderLayer(GradientCheckpointingLayer): """ Decoder layer with standard residual connections. Flow: 1. ActiveNorm(RMSNorm) → LNS(1/√ℓ) → Attention → residual + GPAS 2. ActiveNorm(RMSNorm) → LNS(1/√ℓ) → MLP → Δm 3. h^{ℓ+1} = h̃ + Δm """ def __init__(self, config: NeoLLMConfig, layer_idx: int): super().__init__() self.hidden_size = config.hidden_size self.layer_idx = layer_idx self.use_lns = bool(getattr(config, "use_lns", False)) self.use_gpas = bool(getattr(config, "use_gpas", False)) self.use_siamesenorm = bool(getattr(config, "use_siamesenorm", False)) self.siamese_normalized_input = bool( getattr(config, "siamese_normalized_input", True) ) self.siamese_depth_scaling = bool( getattr(config, "siamese_depth_scaling", True) ) self.siamese_attn_x_scale_init = float( getattr(config, "siamese_attn_x_scale_init", 1.0) ) siamese_stream_scale = ( 1.0 if not self.siamese_depth_scaling else 1.0 / math.sqrt(2.0 * float(layer_idx + 1)) ) # A tensor buffer keeps the depth-dependent value out of Python control # flow during forward. Computing it from ``self.layer_idx`` inside a # Dynamo resume frame specialized that frame once per decoder layer. # Non-persistent preserves checkpoint compatibility: the value is fully # determined by layer_idx and the config on every construction. self.register_buffer( "_siamese_stream_scale_value", torch.tensor(siamese_stream_scale, dtype=torch.float32), persistent=True, ) # Controls only the first pre-attention normalisation applied directly # to the embedding stream. Defaults to True for checkpoint/config # backward compatibility. When False, layer 0 does not instantiate # input_layernorm at all, so the flag removes the corresponding # active norm parameters instead of merely bypassing them in forward. self.use_embedding_input_norm = bool( getattr(config, "use_embedding_input_norm", True) ) self.has_input_layernorm = (not self.use_siamesenorm) and not ( self.layer_idx == 0 and not self.use_embedding_input_norm ) self.self_attn = NeoLLMAttention(config, layer_idx) self.mlp = NeoLLMMLP(config) # JTok/JTok-M operates on the shared MLP update. It is instantiated # only when the corresponding config flag is active; the default # training graph therefore keeps the original residual topology. self.use_jtok = bool(getattr(config, "use_jtok", False)) self.use_jtokm = bool(getattr(config, "use_jtokm", False)) if self.use_jtokm and not self.use_jtok: # ``NeoLLMConfig`` validates this relationship, but decoder layers # are also constructed directly by a few checkpoint conversion and # test paths. Fail explicitly there instead of silently dropping # JTok-M while leaving the rest of the model apparently valid. raise ValueError( "`use_jtokm=True` requires `use_jtok=True`; enable the base " "JTok path before enabling JTok-M." ) self.jtok = ( LeviathanJTok(config, layer_idx) if self.use_jtok else None ) if self.use_siamesenorm: self.input_layernorm = None self.post_attention_layernorm = None # SiameseNorm is RMS-only by config validation. These modules are # constructed only when the Siamese topology is active, so no # inactive RMSNorm pre-norm modules remain in the graph. self.siamese_attn_x_norm = nn.RMSNorm( config.hidden_size, eps=config.rms_norm_eps ) self.siamese_attn_y_norm = nn.RMSNorm( config.hidden_size, eps=config.rms_norm_eps ) self.siamese_mlp_x_norm = nn.RMSNorm( config.hidden_size, eps=config.rms_norm_eps ) self.siamese_mlp_y_norm = nn.RMSNorm( config.hidden_size, eps=config.rms_norm_eps ) self.siamese_attn_input_norm = ( nn.RMSNorm(config.hidden_size, eps=config.rms_norm_eps) if self.siamese_normalized_input else None ) self.siamese_mlp_input_norm = ( nn.RMSNorm(config.hidden_size, eps=config.rms_norm_eps) if self.siamese_normalized_input else None ) self.siamese_attn_x_scale = nn.Parameter( torch.full( (config.hidden_size,), self.siamese_attn_x_scale_init, dtype=torch.float32, ) ) else: self.input_layernorm = ( _make_norm( config.hidden_size, eps=config.rms_norm_eps, ) if self.has_input_layernorm else None ) self.post_attention_layernorm = _make_norm( config.hidden_size, eps=config.rms_norm_eps, ) self.siamese_attn_x_norm = None self.siamese_attn_y_norm = None self.siamese_mlp_x_norm = None self.siamese_mlp_y_norm = None self.siamese_attn_input_norm = None self.siamese_mlp_input_norm = None self.siamese_attn_x_scale = None self.lns_attn = LNS(layer_idx) if self.use_lns else None self.lns_mlp = LNS(layer_idx) if self.use_lns else None self.gpas_attn = GPAS(config.hidden_size) if self.use_gpas else None self.gpas_mlp = GPAS(config.hidden_size) if self.use_gpas else None self.current_layer_fan = None # ── StackMemory / STACKTRANS (Zhang et al., NeurIPS 2025) ──────── # Optional differentiable hidden-state stack inserted between # Transformer layers. The stack is applied before this layer's # attention block, matching the released StackTrans source flow. self.use_stack_memory = getattr(config, "use_stack_memory", False) self.stack_memory = StackMemory(config) if self.use_stack_memory else None # ── Attention Residuals (Kimi Team, 2026) ───────────────────────── # Replaces fixed residual accumulation with learned softmax attention # over preceding layer outputs. Each decoder layer has two learnable # pseudo-queries — one for pre-attention and one for pre-MLP — plus a # shared RMSNorm applied to keys to prevent magnitude-dominated softmax. # Pseudo-queries are initialized to ZERO so AttnRes starts as uniform # average (equivalent to standard residual mean) and training volatility # is avoided. This is critical per the paper's ablation. self.use_attn_res = getattr(config, "use_attn_res", False) if self.use_attn_res: self.attn_res_query_attn = nn.Parameter(torch.zeros(config.hidden_size)) self.attn_res_query_mlp = nn.Parameter(torch.zeros(config.hidden_size)) self.attn_res_norm = nn.RMSNorm(config.hidden_size, eps=config.rms_norm_eps) else: self.attn_res_query_attn = None self.attn_res_query_mlp = None self.attn_res_norm = None # ── LAuReL: Learned Augmented Residual Layer (Menghani et al., ICML 2025) ─ # Generalises the canonical residual connection with learned scalar # weights (RW) and/or a low-rank linear correction (LR). Applied # independently to the attention and MLP sublayers (two residual # junctions per decoder layer). # # LAUREL-RW (use_laurel_rw): # raw scalars α̃, β̃ → softmax([α̃, β̃]) = (α, β) bounded in (0,1) # Residual becomes: α·f(x) + β·x # # LAUREL-LR (use_laurel_lr): # A: nn.Linear(D→r, bias=False) initialised column-orthogonal # B: nn.Linear(r→D, bias=False) initialised to zero # Residual becomes: f(x) + B(A(x)) + x # # LAUREL-RW+LR (both active, paper eq. 5): # Residual becomes: α·f(x) + β·(B(A(x)) + x) # # Mutex with use_attn_res is enforced at config validation time. self.use_laurel = getattr(config, "use_laurel", False) self.use_laurel_rw = getattr(config, "use_laurel_rw", True) self.use_laurel_lr = getattr(config, "use_laurel_lr", False) D = config.hidden_size r = getattr(config, "laurel_lr_rank", 32) if self.use_laurel and self.use_laurel_rw: # Two raw scalars per sublayer; softmax-normalised in forward. # Initialised to [0, 0] → softmax → (0.5, 0.5) at step 0, # matching a standard equal-weight residual as the starting point. self.laurel_rw_attn = nn.Parameter(torch.zeros(2)) # [α̃_attn, β̃_attn] self.laurel_rw_mlp = nn.Parameter(torch.zeros(2)) # [α̃_mlp, β̃_mlp] else: self.laurel_rw_attn = None self.laurel_rw_mlp = None if self.use_laurel and self.use_laurel_lr: # Attention sublayer low-rank matrices. # A: D×r (projects D→r), initialised column-orthogonal per §3.3. # B: r×D (projects r→D), initialised to zero → identity start. self.laurel_lr_A_attn = nn.Linear(D, r, bias=False) self.laurel_lr_B_attn = nn.Linear(r, D, bias=False) # MLP sublayer low-rank matrices (independent capacity). self.laurel_lr_A_mlp = nn.Linear(D, r, bias=False) self.laurel_lr_B_mlp = nn.Linear(r, D, bias=False) # Initialise: B→zero, A→column-orthogonal (paper footnote 2): # A_{i,j} = 1/√(rD) if i mod r == j else 0 for A_mat in (self.laurel_lr_A_attn, self.laurel_lr_A_mlp): nn.init.zeros_(A_mat.weight) for j in range(r): for i in range(D): if i % r == j: A_mat.weight.data[j, i] = 1.0 / (r * D) ** 0.5 for B_mat in (self.laurel_lr_B_attn, self.laurel_lr_B_mlp): nn.init.zeros_(B_mat.weight) else: self.laurel_lr_A_attn = None self.laurel_lr_B_attn = None self.laurel_lr_A_mlp = None self.laurel_lr_B_mlp = None def apply_stack_memory( self, hidden_states: torch.Tensor, stack_memory: Optional[torch.Tensor], stack_memory_mask: Optional[torch.Tensor], token_mask: Optional[torch.Tensor] = None, ) -> Tuple[ torch.Tensor, Optional[torch.Tensor], Optional[torch.Tensor], Optional[torch.Tensor], ]: """ Apply this layer's differentiable StackMemory module before attention. The memory tensors are owned by ``NeoLLMModel.forward`` and threaded across decoder layers, reproducing the StackTrans pattern where a hidden-state stack sits between standard Transformer layers while the attention implementation remains unchanged. """ if (not self.use_stack_memory) or self.stack_memory is None: return hidden_states, stack_memory, stack_memory_mask, None if stack_memory is None or stack_memory_mask is None: raise ValueError( "StackMemory is enabled, but stack_memory/stack_memory_mask " "were not initialized by NeoLLMModel.forward." ) return self.stack_memory( hidden_states, stack_memory, stack_memory_mask, token_mask=token_mask ) def _attn_res( self, sources: list, partial: torch.Tensor, query: torch.Tensor, ) -> torch.Tensor: """ Depth-wise softmax attention over preceding layer outputs. Computes: V = stack(sources + [partial]) [N+1, B, S, D] K = RMSNorm(V) [N+1, B, S, D] logits = query · K [N+1, B, S] weights = softmax(logits, dim=0) [N+1, B, S] h = Σ_n weights_n · V_n [B, S, D] The pseudo-query is shared across positions (per the paper design). RMSNorm on keys prevents layers with large-magnitude outputs from dominating the softmax. Initialized to zero → uniform weights at step 0, reducing to standard residual mean. Args: sources: list of [B, S, D] tensors — completed block summaries or all previous layer outputs (Full AttnRes). partial: [B, S, D] — current intra-block partial sum. query: [D] — learnable pseudo-query for this sublayer. Returns: [B, S, D] — weighted combination of sources + partial. """ all_v = sources + [partial] # list of [B, S, D] V = torch.stack(all_v, dim=0) # [N+1, B, S, D] K = self.attn_res_norm(V) # [N+1, B, S, D] logits = torch.einsum("d,nbsd->nbs", query, K) # [N+1, B, S] weights = torch.softmax(logits, dim=0) # [N+1, B, S] return torch.einsum("nbs,nbsd->bsd", weights, V) # [B, S, D] def _laurel_residual( self, residual: torch.Tensor, delta: torch.Tensor, rw_param: Optional[torch.Tensor], A_mat, B_mat, slot: str = "attn", ) -> torch.Tensor: """ Computes the LAuReL-augmented residual junction (Menghani et al., ICML 2025). Dispatches among three regimes depending on which sub-variants are active: LAUREL-RW only (use_laurel_rw=True, use_laurel_lr=False): α, β = softmax([α̃, β̃]) out = α · delta + β · residual (paper §2.1) LAUREL-LR only (use_laurel_rw=False, use_laurel_lr=True): lr_delta = B(A(residual)) out = delta + lr_delta + residual (paper eq. 3) LAUREL-RW+LR (both active, paper eq. 5): α, β = softmax([α̃, β̃]) lr_delta = B(A(residual)) out = α · delta + β · (lr_delta + residual) In all cases the standard residual identity (out = delta + residual) is recovered at initialisation: RW starts at (α=0.5, β=0.5); LR starts with B=0 so lr_delta=0. Args: residual: [B, S, D] — accumulated residual stream (x_i in the paper). delta: [B, S, D] — sublayer output f(x_i) (attention or MLP). rw_param: Parameter[2] — raw [α̃, β̃] before softmax; None if RW off. A_mat: nn.Linear(D→r, bias=False) — LR down-proj; None if LR off. B_mat: nn.Linear(r→D, bias=False) — LR up-proj; None if LR off. slot: "attn" or "mlp". Returns: [B, S, D] — augmented residual output for the current sublayer. """ has_rw = rw_param is not None has_lr = A_mat is not None # ── LAUREL-LR: low-rank residual correction ──────────────────────── if has_lr: lr_delta = B_mat(A_mat(residual)) # [B, S, D] else: lr_delta = None # ── LAUREL-RW: learned scalar gate ──────────────────────────────── if has_rw: ab = torch.softmax(rw_param.float(), dim=0).to(residual.dtype) alpha = ab[0] beta = ab[1] else: alpha = beta = None # ── Compose output ───────────────────────────────────────────────── if has_rw and has_lr: # LAUREL-RW+LR (paper eq. 5): α·f(x) + β·(BAx + x) return alpha * delta + beta * (lr_delta + residual) elif has_rw: # LAUREL-RW (paper §2.1): α·f(x) + β·x return alpha * delta + beta * residual else: # LAUREL-LR (paper eq. 3): f(x) + BAx + x return delta + lr_delta + residual def _siamese_stream_scale(self, ref: torch.Tensor) -> torch.Tensor: return self._siamese_stream_scale_value.to( device=ref.device, dtype=ref.dtype, ) def forward_siamesenorm( self, x_states: torch.Tensor, y_states: torch.Tensor, position_embeddings: tuple[torch.Tensor, torch.Tensor], stack_memory: Optional[torch.Tensor] = None, stack_memory_mask: Optional[torch.Tensor] = None, stack_token_mask: Optional[torch.Tensor] = None, attention_mask: Optional[torch.Tensor] = None, first_layer_fan: Optional[torch.Tensor] = None, output_attentions: Optional[bool] = False, repo_rope_args: Optional[Tuple[torch.Tensor, float]] = None, position_ids: Optional[torch.LongTensor] = None, jtok_z_tilde: Optional[torch.Tensor] = None, jtok_valid_mask: Optional[torch.Tensor] = None, jtok_compute_aux: bool = False, **kwargs: Unpack[FlashAttentionKwargs], ) -> Tuple: # SiameseNorm keeps two coupled streams with shared Attention/MLP # parameters. All Siamese normalization modules are RMSNorm by # construction; no dynamic-normalization branch remains. # ── Attention shared block ──────────────────────────────────────── x_attn_norm = self.siamese_attn_x_norm(x_states) y_attn_norm = self.siamese_attn_y_norm(y_states) x_scale = self.siamese_attn_x_scale.to( dtype=x_attn_norm.dtype, device=x_attn_norm.device ) h_attn = x_scale * x_attn_norm + y_attn_norm if self.siamese_attn_input_norm is not None: h_attn = self.siamese_attn_input_norm(h_attn) stack_metrics = None if self.use_stack_memory: h_attn, stack_memory, stack_memory_mask, stack_metrics = ( self.apply_stack_memory( h_attn, stack_memory, stack_memory_mask, token_mask=stack_token_mask, ) ) jtok_router_state = h_attn h_lns = self.lns_attn(h_attn) if self.use_lns else h_attn attn_out, attn_weights, self.current_layer_fan = self.self_attn( hidden_states=h_lns, key_value_states=None, attention_mask=attention_mask, position_embeddings=position_embeddings, first_layer_fan=first_layer_fan, repo_rope_args=repo_rope_args, position_ids=position_ids, **kwargs, ) stream_scale = self._siamese_stream_scale(attn_out) x_after_attn = x_states + stream_scale * attn_out y_after_attn = y_states + attn_out if self.use_gpas: x_after_attn = self.gpas_attn(x_after_attn) # ── MLP shared block ────────────────────────────────────────────── x_mlp_norm = self.siamese_mlp_x_norm(x_after_attn) y_mlp_norm = self.siamese_mlp_y_norm(y_after_attn) h_mlp = x_mlp_norm + y_mlp_norm if self.siamese_mlp_input_norm is not None: h_mlp = self.siamese_mlp_input_norm(h_mlp) h_lns2 = self.lns_mlp(h_mlp) if self.use_lns else h_mlp delta_m = self.mlp(h_lns2) jtok_aux_stats = None if self.jtok is not None: if jtok_z_tilde is None: raise ValueError( "Active JTok/JTok-M requires Leviathan `jtok_z_tilde`." ) shared_update, jtok_aux_stats = self.jtok( delta_m, jtok_z_tilde, router_state=jtok_router_state if self.use_jtokm else None, valid_mask=jtok_valid_mask, compute_aux=bool(jtok_compute_aux and self.use_jtokm), ) else: shared_update = delta_m x_next = x_after_attn + stream_scale * shared_update y_next = y_after_attn + shared_update if self.use_gpas: x_next = self.gpas_mlp(x_next) outputs = (x_next, y_next) if self.use_stack_memory: outputs += (stack_memory, stack_memory_mask, stack_metrics) if jtok_aux_stats is not None: outputs += jtok_aux_stats if output_attentions: outputs += (attn_weights,) return outputs def forward( self, hidden_states: torch.Tensor, position_embeddings: tuple[torch.Tensor, torch.Tensor], attention_mask: Optional[torch.Tensor] = None, first_layer_fan: Optional[torch.Tensor] = None, attn_res_sources: Optional[list] = None, attn_res_partial: Optional[torch.Tensor] = None, output_attentions: Optional[bool] = False, repo_rope_args: Optional[Tuple[torch.Tensor, float]] = None, position_ids: Optional[torch.LongTensor] = None, jtok_z_tilde: Optional[torch.Tensor] = None, jtok_valid_mask: Optional[torch.Tensor] = None, jtok_compute_aux: bool = False, **kwargs: Unpack[FlashAttentionKwargs], ) -> Tuple: # ── Snapshot input ──────────────────────────────────────────────── # ── Attention Residuals: compute pre-attention input ────────────── # When active, the input to the attention sublayer is no longer the # raw hidden_states (accumulated residual) but a softmax-weighted # combination of all previous layer outputs (or block summaries). # attn_res_partial carries the intra-block standard residual that # connects the attention and MLP sublayers within this layer. # When inactive, flow is identical to the original. if ( self.use_attn_res and attn_res_sources is not None and attn_res_partial is not None ): h_attn = self._attn_res( attn_res_sources, attn_res_partial, self.attn_res_query_attn, ) residual_attn = attn_res_partial else: h_attn = hidden_states residual_attn = hidden_states # ── Attention block ─────────────────────────────────────────────── # Optional embedding input normalisation: # layer_idx == 0: h_attn is the embedding stream entering the first # attention block, so the flag can expose raw # embedding magnitudes to attention. # layer_idx > 0: keep the standard pre-norm Transformer flow exactly # as before. if self.input_layernorm is None: h_normed = h_attn else: h_normed = _apply_norm(self.input_layernorm, h_attn) # JTok-M follows the source design and routes from the fused # pre-attention contextual state, before the attention output is added. jtok_router_state = h_normed h_lns = self.lns_attn(h_normed) if self.use_lns else h_normed hidden_states, attn_weights, self.current_layer_fan = self.self_attn( hidden_states=h_lns, key_value_states=None, attention_mask=attention_mask, position_embeddings=position_embeddings, first_layer_fan=first_layer_fan, repo_rope_args=repo_rope_args, position_ids=position_ids, **kwargs, ) # ── Residual junction (attention sublayer) ──────────────────────── if self.use_laurel: attn_aug = self._laurel_residual( residual_attn, hidden_states, self.laurel_rw_attn, self.laurel_lr_A_attn, self.laurel_lr_B_attn, slot="attn", ) else: attn_aug = residual_attn + hidden_states h_tilde = self.gpas_attn(attn_aug) if self.use_gpas else attn_aug # ── Attention Residuals: compute pre-MLP input ──────────────────── # After attention, the partial sum is updated with h_tilde. # The pre-MLP AttnRes attends over the same sources but with h_tilde # as the current partial — capturing the within-layer attention output. if self.use_attn_res and attn_res_sources is not None: h_mlp = self._attn_res( attn_res_sources, h_tilde, self.attn_res_query_mlp, ) residual_mlp = h_tilde else: h_mlp = h_tilde residual_mlp = h_tilde # ── MLP block ───────────────────────────────────────────────────── h_normed2 = _apply_norm(self.post_attention_layernorm, h_mlp) h_lns2 = self.lns_mlp(h_normed2) if self.use_lns else h_normed2 delta_m = self.mlp(h_lns2) jtok_aux_stats = None if self.jtok is not None: if jtok_z_tilde is None: raise ValueError( "Active JTok/JTok-M requires Leviathan `jtok_z_tilde`." ) delta_m, jtok_aux_stats = self.jtok( delta_m, jtok_z_tilde, router_state=jtok_router_state if self.use_jtokm else None, valid_mask=jtok_valid_mask, compute_aux=bool(jtok_compute_aux and self.use_jtokm), ) # ── Residual junction (MLP sublayer) ───────────────────────────── # LAuReL augments the base MLP residual (residual_mlp + delta_m). if self.use_laurel: mlp_aug = self._laurel_residual( residual_mlp, delta_m, self.laurel_rw_mlp, self.laurel_lr_A_mlp, self.laurel_lr_B_mlp, slot="mlp", ) else: mlp_aug = residual_mlp + delta_m hidden_states = self.gpas_mlp(mlp_aug) if self.use_gpas else mlp_aug outputs = (hidden_states,) if jtok_aux_stats is not None: outputs += jtok_aux_stats if output_attentions: outputs += (attn_weights,) return outputs class SpellingBeeEmbedding(nn.Module): """ Spelling Bee Embeddings (Rabe et al., 2026, arXiv:2601.18030). Augments token embeddings with character-level information derived from the UTF-8 byte sequence of each token. The spelling bee embedding is the mean of the standard token embedding and a character-level summary: e_bee(t) = 0.5 * (e_tok(t) + e_chars(t)) e_chars(t) = (1 / MAX_BYTES) * Σ_{i=0}^{15} RoPE(e_byte[b_i], i) The active path uses the fixed 16-position mean from the reference implementation. A theoretical alternative is to mask padded positions and normalize only the valid bytes with ``1 / sqrt(|t|)``; that variant is documented in the configuration but is not active here. Key design decisions vs. a naïve per-occurrence implementation: 1. **Vocab-level computation** — e_chars is built over the full vocabulary once per forward (shape [V, d]), then gathered by token_ids. A naïve implementation would compute [B*S, 16, d] per step, repeating identical work for every occurrence of a frequent token. This approach reduces the dominant intermediate from O(B·S·16·d) to O(V·16·d), where V ≪ B·S in practice for most batches. 2. **Static [256, 16, d] rope_bytes table** — RoPE is applied once over all 256 possible byte values at all 16 positions, producing a table with fully static shapes. torch.compile / max_autotune can fuse the construction of this table (two elementwise ops + concat over fixed dims) into a single kernel. Token-level e_chars is then a gather + sum over this table, also fully static. 3. **Fixed-length mean** — every token contributes exactly 16 position-aware byte slots, including the learned null-byte padding, and the sum is divided by ``MAX_BYTES``. This keeps the active path faithful to the reference experiment and avoids a per-token scale factor. Compatible with both the standard embed_tokens path and the LeviathanGenerator path. **Inference cost: zero overhead after baking.** Call ``bake_inference_table(token_embeds_weight)`` once after training to collapse the SBE into a single embedding table indistinguishable from a standard nn.Embedding lookup. **Setup: call ``set_byte_table(tokenizer)`` once after model init** (and before any .to(device) / FP8 conversion) before training. The byte table token_bytes is a persistent buffer saved in checkpoints. References: Rabe, Clymo & Dong (2026). "Spelling Bee Embeddings for Language Modeling." arXiv:2601.18030. """ MAX_BYTES: int = 16 def __init__(self, config: "NeoLLMConfig"): super().__init__() d = config.hidden_size base = getattr(config, "rope_theta", 10000.0) # Guardado para poder recomputar los buffers RoPE en _reset_rope_buffers. self._rope_base = base # 256 × d byte embedding lookup (one per UTF-8 byte value 0..255). self.byte_emb = nn.Embedding(256, d) # ── Persistent buffers (saved in checkpoints) ───────────────────── # token_bytes [vocab_size, MAX_BYTES]: UTF-8 byte values per token, # padded with 0x00 up to MAX_BYTES positions. self.register_buffer( "token_bytes", torch.zeros(config.vocab_size, self.MAX_BYTES, dtype=torch.long), persistent=True, ) def _load_from_state_dict( self, state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs, ): """ Sobreescribe la carga de state_dict para eliminar los tres buffers non-persistent (intra_cos, intra_sin, pos_idx) antes de aplicar el state_dict, evitando que versiones anteriores del checkpoint (donde eran persistent=True) sobreescriban los valores correctos calculados en __init__ con valores corruptos del safetensors. from_pretrained de HuggingFace bypasea _register_load_state_dict_pre_hook y carga directamente por nombre, por lo que este override es necesario. """ for key in ("intra_cos", "intra_sin", "pos_idx"): state_dict.pop(prefix + key, None) super()._load_from_state_dict( state_dict, prefix, local_metadata, strict, missing_keys, unexpected_keys, error_msgs, ) def set_byte_table(self, tokenizer) -> None: """ Precompute the UTF-8 byte table from a tokenizer. Must be called **once** after model instantiation and **before** ``.to(device)`` / FP8 conversion so the buffers land on the correct device after those transforms. Both buffers are persistent and will be saved/restored from checkpoints automatically. Args: tokenizer: Any tokenizer with ``decode([token_id]) -> str``. """ vocab_size = self.token_bytes.shape[0] byte_ids = torch.zeros(vocab_size, self.MAX_BYTES, dtype=torch.long) for token_id in range(vocab_size): # Decode through the active tokenizer. convert_ids_to_tokens() # returns an internal BPE/SentencePiece symbol (for example # ``Ġapple`` or ``Ċ``), not the text surface (`` apple`` or ``\n``) # whose UTF-8 spelling Spelling Bee is meant to represent. try: try: token_str = tokenizer.decode( [token_id], skip_special_tokens=False, clean_up_tokenization_spaces=False, ) except TypeError: # Keep compatibility with tokenizer implementations that # do not expose clean_up_tokenization_spaces. token_str = tokenizer.decode( [token_id], skip_special_tokens=False ) raw = token_str.encode("utf-8") except Exception: raw = b"\x00" n = min(len(raw), self.MAX_BYTES) for i in range(n): byte_ids[token_id, i] = raw[i] self.token_bytes.copy_(byte_ids.to(self.token_bytes.device)) # ── Core helpers ────────────────────────────────────────────────────────── def _build_rope_bytes(self) -> torch.Tensor: """ Build the static [256, MAX_BYTES, d] RoPE-encoded byte table. For each of the 256 possible byte values and each of the MAX_BYTES intra-token positions, applies RoPE rotation using the current byte_emb.weight. All shapes are fully static, so torch.compile can fuse this into a single kernel. cos/sin se computan inline aquí en vez de usar buffers registrados. Con device_map + accelerate, from_pretrained materializa tensores del safetensors directamente — non-persistent buffers que no están en el checkpoint quedan como memoria sin inicializar. Computar inline elimina esa dependencia sin overhead apreciable ([16, d//2] es ínfimo). Returns: rope_bytes [256, MAX_BYTES, d] """ w = self.byte_emb.weight # [256, d] half = w.shape[-1] // 2 w1 = w[:, :half].unsqueeze(1) # [256, 1, half] w2 = w[:, half:].unsqueeze(1) # [256, 1, half] # Computar RoPE inline — formas estáticas, torch.compile lo fusiona. theta = 1.0 / ( self._rope_base ** ( torch.arange(0, half, dtype=torch.float32, device=w.device) * 2.0 / (half * 2) ) ) pos = torch.arange(self.MAX_BYTES, dtype=torch.float32, device=w.device) freqs = torch.outer(pos, theta) # [MAX_BYTES, half] cos = freqs.cos().to(w.dtype).unsqueeze(0) # [1, MAX_BYTES, half] sin = freqs.sin().to(w.dtype).unsqueeze(0) # [1, MAX_BYTES, half] return torch.cat( [w1 * cos - w2 * sin, w1 * sin + w2 * cos], dim=-1, ) # [256, MAX_BYTES, d] # ── Forward ─────────────────────────────────────────────────────────────── def forward( self, token_ids: torch.Tensor, # [B, S] or [N] token_embeds: torch.Tensor, # [B, S, d] or [N, d] *, meap_mask_embedding: Optional[torch.Tensor] = None, meap_mask_token_id: Optional[int] = None, ) -> torch.Tensor: """ Args: token_ids: integer token indices to look up byte sequences. token_embeds: embeddings from embed_tokens or LeviathanGenerator. Returns: Spelling bee embeddings — same shape as token_embeds. """ # ── Step 1: build rope_bytes over 256 byte types × 16 positions ─── # Shape [256, MAX_BYTES, d] — fully static, one kernel via compile. rope_bytes = self._build_rope_bytes() # [256, MAX_BYTES, d] # ── Step 2: gather bytes only for token occurrences in this batch ─── # Building [V, MAX_BYTES, d] here wastes more than 1 GiB for a 64k # vocabulary and retains that full graph for backward. Indexing the # byte table first is mathematically identical for the requested token # ids while keeping the temporary proportional to B×S instead of V. token_byte_rows = self.token_bytes[token_ids] # [..., MAX_BYTES] e_chars = rope_bytes[ token_byte_rows, torch.arange(self.MAX_BYTES, device=rope_bytes.device), ].mean(-2) # [..., d], fixed 16-position mean # ── Step 3: mean with token embeddings ────────────────────────────── output = (token_embeds + e_chars) * 0.5 has_meap = ( meap_mask_embedding is not None or meap_mask_token_id is not None ) if not has_meap: return output if meap_mask_embedding is None or meap_mask_token_id is None: raise ValueError( "meap_mask_embedding and meap_mask_token_id must be provided " "together" ) return torch.where( token_ids.eq(int(meap_mask_token_id)).unsqueeze(-1), meap_mask_embedding.to(device=output.device, dtype=output.dtype), output, ) # ── Inference utility ───────────────────────────────────────────────────── @torch.no_grad() def bake_inference_table( self, token_emb_weight: torch.Tensor, ) -> torch.Tensor: """ Collapse SBE into a single [vocab_size, d] embedding table. After baking, the SBE computation is indistinguishable from a standard nn.Embedding lookup — zero additional overhead at inference time. Args: token_emb_weight: [vocab_size, d] — weight matrix of embed_tokens or the equivalent table (e.g. after Leviathan). Returns: [vocab_size, d] — baked spelling bee embedding table. Usage:: baked = model.model.spelling_bee.bake_inference_table( model.model.embed_tokens.weight ) model.model.embed_tokens.weight.copy_(baked) # Optionally free byte_emb parameters: # del model.model.spelling_bee """ rope_bytes = self._build_rope_bytes() # [256, MAX_BYTES, d] e_chars_vocab = rope_bytes[ self.token_bytes, torch.arange(self.MAX_BYTES, device=rope_bytes.device), ].mean(1) # [V, d], fixed 16-position mean return (token_emb_weight + e_chars_vocab) * 0.5 class NeoLLMPreTrainedModel(PreTrainedModel): """ Base class with custom weight initialization for all NeoLLM components. LeviathanGenerator (LEV layer, paper-faithful per-head architecture): - codebooks: normal(0, initializer_range) - head_proj_weight: normal(0, initializer_range) — per-head seed projections W_seed,l - head_norm_weight/bias: weight=1, bias=0 — default LayerNorm init - head_spline_delta: normal(mean=0.0, std=0.1). The effective coefficient is (1 + delta), matching the Leviathan reference parameterization and keeping the product across d_seed dimensions near 1 at step 0. - head_out_weight: normal(0, initializer_range / sqrt(num_modes)) — scaled by 1/√h so the sum of h head outputs starts with the same variance as a single head projection. NITPTemporalTransition: - context_gate_proj/context_up_proj/condition_modulation_proj: Xavier-uniform so both the contextual proposal and the detached NITP condition have non-degenerate variance at step 0. - down_proj: normal(0, 0.01) — keeps the residual transition close to identity initially without blocking gradients into the centered condition gate. - temporal_step_bias: zeros — all horizons start with the same phase. During the forward pass the bias table is projected to zero mean over the horizon, so it can represent only relative phase differences, not one shared phase offset. - temporal_step_gain_logits: zeros — 2*sigmoid(0)=1, so every horizon starts with exact unit residual gain. - context_norm/condition_norm: weight=1. NeoLLMAttention (Affine-Scaled Attention): - alpha_proj: normal(0, 0.02) — near-zero so linear_clipping(≈0) ≈ 0.5 at init, giving a mild ~0.5× scaling of softmax weights per head rather than collapsing to 0 or 1. - alpha_ma: zeros — running EMA starts at 0, β starts as −α/N ≈ small negative offset; model quickly learns to adjust both. REPOModule (Context Re-Positioning): - W_g, W_c, W_z: default normal init from parent _init_weights. No special initialization required — the SwiGLU sub-layer starts near-zero, so z_i ≈ 0 for all tokens at step 0, which is equivalent to constant position assignment (NoPE-like). The model quickly learns to differentiate positions as needed. """ config: NeoLLMConfig base_model_prefix = "model" supports_gradient_checkpointing = True _no_split_modules = ["NeoLLMDecoderLayer"] _supports_attention_backend = True _supports_flash_attn = True _supports_flash_attn_2 = True _supports_sdpa = True _is_stateful = True @classmethod def from_pretrained(cls, *args, **kwargs): """Restore derived buffers and support legacy ``knot_grid`` loading. Current checkpoints persist this buffer and must retain its exact saved values. The previous compatibility path regenerated the grid after every successful load, silently overwriting a valid checkpoint (and BF16 direct-linspace rounding can differ from FP32-then-BF16 rounding). """ return_loading_info = bool(kwargs.pop("output_loading_info", False)) requested_dtype = kwargs.get("dtype", kwargs.get("torch_dtype")) model, loading_info = super().from_pretrained( *args, output_loading_info=True, **kwargs, ) missing_keys = set(loading_info.get("missing_keys", ())) # LNS.scale is derived, non-persistent state. Low-memory loaders can # materialize its meta storage without running the original constructor # initializer. Rebuild the constant; never infer validity from finiteness. for module in model.modules(): if isinstance(module, LNS): scale_device = module.scale.device if scale_device.type == "meta": scale_device = model.device if scale_device.type == "meta": scale_device = torch.device("cpu") scale_dtype = ( requested_dtype if isinstance(requested_dtype, torch.dtype) and requested_dtype.is_floating_point else module.scale.dtype ) module._buffers["scale"] = torch.tensor( 1.0 / math.sqrt(module.layer_idx), device=scale_device, dtype=torch.float32, ).to(dtype=scale_dtype) for module_name, module in model.named_modules(): if not isinstance(module, LeviathanGenerator): continue checkpoint_key = ( f"{module_name}.knot_grid" if module_name else "knot_grid" ) needs_legacy_initialization = ( module.knot_grid.device.type == "meta" or checkpoint_key in missing_keys ) if not needs_legacy_initialization: continue if module.knot_grid.device.type == "meta": module._buffers["knot_grid"] = torch.linspace( 0.0, 1.0, module.num_knots, device=torch.device("cpu"), dtype=torch.float32, ) else: with torch.no_grad(): # Match training construction: create the fixed grid in # FP32 first, then let the explicit cast perform rounding. module.knot_grid.copy_( torch.linspace( 0.0, 1.0, module.num_knots, device=module.knot_grid.device, dtype=torch.float32, ).to(dtype=module.knot_grid.dtype) ) # These SiameseNorm tensors are constructed explicitly in FP32 to make # direct ``NeoLLMForCausalLM(config)`` initialization stable. Training # then casts them with ``model.to(BF16)``. Transformers' low-memory # loader honors the explicit constructor dtype instead of ``dtype=`` # when copying a checkpoint, so without this normalization a BF16 Hub # checkpoint silently resumes those 12 trainable vectors in FP32. if ( isinstance(requested_dtype, torch.dtype) and requested_dtype.is_floating_point ): for module in model.modules(): siamese_scale = getattr(module, "siamese_attn_x_scale", None) if siamese_scale is not None: siamese_scale.data = siamese_scale.data.to( dtype=requested_dtype ) stream_scale = getattr( module, "_siamese_stream_scale_value", None, ) if stream_scale is not None: module._buffers["_siamese_stream_scale_value"] = ( stream_scale.to(dtype=requested_dtype) ) if return_loading_info: return model, loading_info return model def _init_weights(self, module): super()._init_weights(module) if isinstance(module, NITPTemporalTransition): # A centered tanh gate has no constant 0.5 path. Xavier # initialization gives the detached NITP condition measurable # variation from step 0, while a small non-zero down projection # keeps the full residual transition close to identity without # suppressing gradients into the condition path. nn.init.xavier_uniform_(module.context_gate_proj.weight) nn.init.xavier_uniform_(module.context_up_proj.weight) nn.init.xavier_uniform_(module.condition_modulation_proj.weight) nn.init.normal_(module.down_proj.weight, mean=0.0, std=0.01) nn.init.zeros_(module.temporal_step_bias) nn.init.zeros_(module.temporal_step_gain_logits) module.context_norm.weight.data.fill_(1.0) module.condition_norm.weight.data.fill_(1.0) if isinstance(module, NeoLLMAttention): if getattr(module, "use_fan_residual", False): if hasattr(module, "lambda_1") and module.lambda_1 is not None: module.lambda_1.data.fill_(0.5) if hasattr(module, "lambda_2") and module.lambda_2 is not None: module.lambda_2.data.fill_(0.5) if hasattr(module, "mea_key_mix") and module.mea_key_mix is not None: # Identity initialization: at step 0 MEA behaves as standard attention # and all matrix entries receive gradient immediately from the first step. # For square matrices (normal training case) this is exact identity. # For rectangular matrices (KV compression, h' Tuple: output_hidden_states = ( output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states ) output_attentions = ( output_attentions if output_attentions is not None else self.config.output_attentions ) output_tweo_activations = bool(output_tweo_activations) if return_dict is None: cfg_dict = vars(self.config) return_dict = cfg_dict.get( "return_dict", cfg_dict.get("use_return_dict", True), ) if (input_ids is None) ^ (inputs_embeds is not None): raise ValueError("Specify exactly one of input_ids or inputs_embeds") use_jtok = bool(getattr(self.config, "use_jtok", False)) use_jtokm = bool(getattr(self.config, "use_jtokm", False)) if use_jtok and input_ids is None: raise ValueError( "Active JTok/JTok-M requires `input_ids`; precomputed " "`inputs_embeds` do not carry the Leviathan token coordinate." ) jtok_z_tilde = None # Diagnostic validity is separate from attention/loss/JTok masks. # In particular, EOS may also be the padding ID; never infer padding # from IDs. Exclude MEAP replacements even if SBE owns the epilogue. dynamics_kwargs = {} if (bool(getattr(self.config, "dynamics_metrics_enabled", False)) and self.training and torch.is_grad_enabled() and input_ids is not None): if attention_mask is not None: if attention_mask.ndim != 2 or attention_mask.shape != input_ids.shape: raise ValueError("Dynamics collection requires an explicit 2D attention validity mask") dynamics_mask = attention_mask.to(dtype=torch.bool) else: dynamics_mask = torch.ones_like(input_ids, dtype=torch.bool) if meap_active: dynamics_mask = dynamics_mask & input_ids.ne(int(self.config.meap_mask_token_id)) dynamics_kwargs = {"dynamics_valid_mask": dynamics_mask} # ── Embedding stage ──────────────────────────────────────────────── meap_mask_embedding = None meap_mask_token_id = None if meap_active: if input_ids is None: raise ValueError("Active MEAP requires `input_ids`.") if self.meap_mask_embedding is None: raise ValueError( "Active MEAP requires a configured `meap_mask_token_id`." ) meap_mask_embedding = self.meap_mask_embedding meap_mask_token_id = int(self.config.meap_mask_token_id) meap_consumed = False if inputs_embeds is None: if self.config.use_token_generator: # Without a later Spelling Bee epilogue, Leviathan owns the # final embedding write and selects the MEAP vector inside its # GEMM store. If SBE is active, its epilogue must own the final # selection so byte features cannot alter masked positions. fuse_meap_in_leviathan = meap_active and self.spelling_bee is None if fuse_meap_in_leviathan: if use_jtok: generator_result = self.token_generator( input_ids, meap_mask_embedding=meap_mask_embedding, meap_mask_token_id=meap_mask_token_id, return_jtok_geometry=True, **dynamics_kwargs, ) else: # Keep the pre-JTok call signature for custom or # checkpoint-restored token generators. generator_result = self.token_generator( input_ids, meap_mask_embedding=meap_mask_embedding, meap_mask_token_id=meap_mask_token_id, **dynamics_kwargs, ) else: if use_jtok: generator_result = self.token_generator( input_ids, return_jtok_geometry=True, **dynamics_kwargs, ) else: generator_result = self.token_generator(input_ids, **dynamics_kwargs) if use_jtok: inputs_embeds, jtok_z_tilde = generator_result else: inputs_embeds = generator_result meap_consumed = fuse_meap_in_leviathan else: inputs_embeds = self.embed_tokens(input_ids) # ── Spelling Bee Embeddings (applied post-embedding, pre-decoder) ────── # input_ids may be None when inputs_embeds was passed directly by the # caller; in that case SBE cannot run (no token_ids available) and is # silently skipped — consistent with the standard embedding bypass path. if self.spelling_bee is not None and input_ids is not None: inputs_embeds = self.spelling_bee( input_ids, inputs_embeds, meap_mask_embedding=( meap_mask_embedding if meap_active else None ), meap_mask_token_id=(meap_mask_token_id if meap_active else None), ) meap_consumed = meap_active # This must be the final operation in the input-embedding pipeline. # The clean/inference path is unchanged and does not depend on CCE. if meap_active and not meap_consumed: meap_positions = input_ids.eq(meap_mask_token_id).unsqueeze(-1) mask_embedding = meap_mask_embedding.to( device=inputs_embeds.device, dtype=inputs_embeds.dtype, ).view(1, 1, -1) inputs_embeds = torch.where( meap_positions, mask_embedding, inputs_embeds, ) # The valid mask serves two separate purposes: JTok/JTok-M must not # modulate padding or MEAP-masked positions, and JTok-M must exclude # those positions from its load-balancing counts. The original # attention mask remains unchanged for positions/loss computation. if use_jtok: jtok_valid_mask = ( attention_mask.to( device=input_ids.device, dtype=torch.bool, ) if attention_mask is not None else torch.ones_like(input_ids, dtype=torch.bool) ) if meap_active: jtok_valid_mask = jtok_valid_mask & input_ids.ne( int(meap_mask_token_id) ) else: jtok_valid_mask = None if position_ids is None: if attention_mask is not None: position_ids = attention_mask.to(dtype=torch.long).cumsum(-1) - 1 position_ids = position_ids.masked_fill(attention_mask == 0, 0) else: position_ids = torch.arange( 0, inputs_embeds.shape[1], device=inputs_embeds.device ).unsqueeze(0) # ── Compile-stable FlashAttention training mask contract ───────────── # The Transformers FA2/FA3 wrapper treats a 2-D padding mask as a # variable-length problem: it branches between ``None`` and a tensor, # calls ``nonzero`` to unpad it, and specializes on both the resulting # token count and ``layer_idx``. With batches containing different # amounts of padding, Dynamo therefore recompiles layer by layer and # retains several compiled graphs, which can *increase* VRAM precisely # when IHA is disabled. IHA's seq-expand kernel does not use that # wrapper, so it did not expose the same behavior. # # Training in this project has a strict dense, causal, right-padding # contract (no packed sequences and no interior holes). Under that # contract, valid queries occur before every padded key and causal # attention cannot see those future padded positions. Outputs for # padded queries are discarded by the masked objectives. Passing # ``None`` to standard decoder attention backends is thus equivalent on # trainable tokens while keeping one fixed dense kernel topology. # # IMPORTANT: the original ``attention_mask`` remains intact above and # outside the decoder: it still builds ``position_ids`` and is still # consumed by CCE/TWEO/NITP/temporal/NextLat and label masking. Do not # broaden this branch to packed/left-padded data. Validate that kind # of input at the data-pipeline boundary rather than with a tensor-value # check here, because a runtime ``all()``/``nonzero()`` would recreate # the graph breaks this branch is intended to remove. Evaluation, # generation and non-FlashAttention backends retain standard mask # construction below. IHA receives the original 2-D mask only as a # static dense/padded sentinel for its packing detector; its direct # kernel intentionally does not unpad fixed-length right-padded batches. use_dense_right_padded_flash_train = ( self.training and self.config._attn_implementation in {"flash_attention_2", "flash_attention_3"} ) if use_dense_right_padded_flash_train: causal_mask = ( attention_mask if bool(getattr(self.config, "use_iha", False)) else None ) else: causal_mask = create_causal_mask( config=self.config, inputs_embeds=inputs_embeds, attention_mask=attention_mask, past_key_values=None, position_ids=position_ids, ) hidden_states = inputs_embeds use_siamesenorm = bool(getattr(self.config, "use_siamesenorm", False)) siamese_y_states = inputs_embeds if use_siamesenorm else None all_hidden_states = () if output_hidden_states else None all_attentions = () if output_attentions else None tweo_activations = () if output_tweo_activations else None collect_jtokm_aux = bool( use_jtokm and self.training and torch.is_grad_enabled() ) all_jtokm_aux_stats = () if collect_jtokm_aux else None # ── StackMemory state ───────────────────────────────────────────── # Llama v9_1 geometry: every token starts with its own empty stack and # retains that token-local [H,K,r] state across decoder layers. This # removes the old ``new_stack[:, -1]`` collapse/broadcast path. use_stack_memory = getattr(self.config, "use_stack_memory", False) if use_stack_memory: batch_size, seq_len = hidden_states.shape[:2] stack_memory = ( self.memory.detach() .to(dtype=hidden_states.dtype) .expand(batch_size, seq_len, -1, -1, -1) ) # Keep the differentiable StackMemory occupancy/control state in # FP32 even when the live model parameters/activations are BF16. stack_memory_mask = ( self.memory_mask.detach() .to(device=hidden_states.device, dtype=torch.float32) .expand(batch_size, seq_len, -1, -1) ) else: stack_memory = None stack_memory_mask = None all_stack_metrics = () if use_stack_memory else None position_embeddings = self.rotary_emb(hidden_states, position_ids) self.first_layer_fan = ( None if getattr(self.config, "use_fan_residual", True) else False ) # ── REPO: pass inv_freq by reference at forward time ────────────────── # rotary_emb.inv_freq is already on the correct device (managed by # NeoLLMRotaryEmbedding as a buffer) — no .to(), no DeviceCopy op. # Computed once here and passed through the decoder layer chain so # NeoLLMAttention never needs to store it as a buffer itself, avoiding # the meta-tensor issue that occurs when lm_eval calls .to(device). # # Extended for IHA seq-expand (P>1): when use_iha=True and P>1, the # attention layer also needs inv_freq to build interleaved RoPE positions # (s·P + p) for the expanded virtual sequence — even on layers that do not # use REPO. Passing repo_rope_args unconditionally in this case incurs # only a small reference pass; it is ignored by layers where neither REPO # nor IHA seq-expand is active. _iha_needs_inv_freq = ( getattr(self.config, "use_iha", False) and int(getattr(self.config, "iha_num_pseudo_heads", 1)) > 1 ) _repo_grape_needs_inv_freq = bool(getattr(self.config, "use_repo_grape", False)) repo_rope_args = ( (self.rotary_emb.inv_freq, self.rotary_emb.attention_scaling) if ( getattr(self.config, "use_repo", False) or _repo_grape_needs_inv_freq or _iha_needs_inv_freq ) else None ) # ── Attention Residuals state ────────────────────────────────────── # Full AttnRes (attn_res_num_blocks=0): sources grows by one entry per # decoder layer — all previous outputs are kept, max N=num_layers+1. # Block AttnRes (attn_res_num_blocks>0): sources grows by one entry per # block boundary — at most num_blocks+1 entries, far less memory. # In both modes, attn_res_partial is the current intra-block accumulated # hidden state that connects the attn and MLP sublayers and flows between # decoder layers within a block. use_attn_res = getattr(self.config, "use_attn_res", False) attn_res_sources = None attn_res_partial = None if use_attn_res: attn_res_sources = [hidden_states] # b_0 = token embedding attn_res_partial = hidden_states # initial partial sum num_blocks = getattr(self.config, "attn_res_num_blocks", 0) block_size = ( max(self.config.num_hidden_layers // num_blocks, 1) if num_blocks > 0 else 1 # Full AttnRes: every layer is its own "block" ) for layer_idx, decoder_layer in enumerate( self.layers[: self.config.num_hidden_layers] ): if output_hidden_states: all_hidden_states = all_hidden_states + (hidden_states,) # ── Block AttnRes: boundary handling ────────────────────────── # At each block boundary (excluding layer 0): append the current # partial sum to sources as a completed block summary, then reset # partial to None so the new block builds from scratch — matching # the paper's pseudocode exactly. # For Full AttnRes (block_size=1): every layer is a boundary, so # partial is appended and reset after every layer. The partial is # re-seeded from the previous hidden_states below. if use_attn_res and layer_idx > 0 and layer_idx % block_size == 0: attn_res_sources = attn_res_sources + [attn_res_partial] attn_res_partial = hidden_states # start new block from current output if use_stack_memory and not use_siamesenorm: hidden_states, stack_memory, stack_memory_mask, stack_layer_metrics = ( decoder_layer.apply_stack_memory( hidden_states, stack_memory, stack_memory_mask, token_mask=attention_mask, ) ) all_stack_metrics = all_stack_metrics + (stack_layer_metrics,) if use_siamesenorm: layer_outputs = decoder_layer.forward_siamesenorm( hidden_states, siamese_y_states, position_embeddings=position_embeddings, stack_memory=stack_memory, stack_memory_mask=stack_memory_mask, stack_token_mask=attention_mask, attention_mask=causal_mask, first_layer_fan=self.first_layer_fan, output_attentions=output_attentions, repo_rope_args=repo_rope_args, position_ids=position_ids, jtok_z_tilde=jtok_z_tilde, jtok_valid_mask=jtok_valid_mask, jtok_compute_aux=collect_jtokm_aux, **kwargs, ) hidden_states = layer_outputs[0] siamese_y_states = layer_outputs[1] if use_stack_memory: stack_memory = layer_outputs[2] stack_memory_mask = layer_outputs[3] stack_layer_metrics = layer_outputs[4] all_stack_metrics = all_stack_metrics + (stack_layer_metrics,) extras_start = 5 else: extras_start = 2 else: layer_outputs = decoder_layer( hidden_states, position_embeddings=position_embeddings, attention_mask=causal_mask, first_layer_fan=self.first_layer_fan, attn_res_sources=attn_res_sources, attn_res_partial=attn_res_partial if use_attn_res else None, output_attentions=output_attentions, repo_rope_args=repo_rope_args, position_ids=position_ids, jtok_z_tilde=jtok_z_tilde, jtok_valid_mask=jtok_valid_mask, jtok_compute_aux=collect_jtokm_aux, **kwargs, ) hidden_states = layer_outputs[0] extras_start = 1 # Update AttnRes partial sum — the new partial is the layer output if use_attn_res: attn_res_partial = hidden_states if collect_jtokm_aux: all_jtokm_aux_stats = all_jtokm_aux_stats + ( ( layer_outputs[extras_start], layer_outputs[extras_start + 1], layer_outputs[extras_start + 2], ), ) extras_start += 3 if output_attentions: all_attentions = all_attentions + (layer_outputs[extras_start],) extras_start += 1 if output_tweo_activations: tweo_activations = tweo_activations + (hidden_states,) if ( getattr(self.config, "use_fan_residual", True) and self.first_layer_fan is None and hasattr(decoder_layer, "current_layer_fan") ): self.first_layer_fan = decoder_layer.current_layer_fan if use_siamesenorm: hidden_states = self.siamese_final_norm( self.siamese_x_final_norm(hidden_states) + self.siamese_y_final_norm(siamese_y_states) ) else: hidden_states = self.norm(hidden_states) if output_hidden_states: all_hidden_states = all_hidden_states + (hidden_states,) stack_metrics = ( torch.stack(all_stack_metrics, dim=0) if all_stack_metrics is not None and len(all_stack_metrics) > 0 else None ) if not return_dict: base_outputs = tuple( v for v in [hidden_states, None, all_hidden_states, all_attentions] if v is not None ) if stack_metrics is not None: base_outputs = base_outputs + (stack_metrics,) if all_jtokm_aux_stats is not None: base_outputs = base_outputs + (all_jtokm_aux_stats,) if output_tweo_activations: return base_outputs + (tweo_activations,) return base_outputs outputs = NeoLLMBaseModelOutputWithPast( last_hidden_state=hidden_states, past_key_values=None, hidden_states=all_hidden_states, attentions=all_attentions, jtokm_aux_stats=all_jtokm_aux_stats, stack_metrics=stack_metrics, ) if output_tweo_activations: return outputs, tweo_activations return outputs def _prepare_lm_labels(labels, device, pad_token_id=None): """Move labels to device and normalize padding to the CE ignore index.""" processed_labels = labels.to(device) if pad_token_id is not None: processed_labels = torch.where( processed_labels == pad_token_id, torch.tensor( -100, dtype=processed_labels.dtype, device=processed_labels.device ), processed_labels, ) return processed_labels def _resolve_liger_accum_dtype(accum_dtype): """Map a serializable config value to the torch.dtype expected by Liger.""" if accum_dtype is None: return None if isinstance(accum_dtype, torch.dtype): return accum_dtype value = str(accum_dtype).strip().lower() if value in {"", "none", "null", "auto", "default", "original"}: return None if value in {"float32", "fp32", "float", "torch.float32"}: return torch.float32 if value in {"bfloat16", "bf16", "torch.bfloat16"}: return torch.bfloat16 if value in {"float16", "fp16", "half", "torch.float16"}: return torch.float16 raise ValueError( "`liger_loss_accum_dtype` must be one of None/'none', " "'float32'/'fp32', 'bfloat16'/'bf16', or 'float16'/'fp16'. " f"Got {accum_dtype!r}." ) _EXTENDED_CCE_INSTALL = ( "cut-cross-entropy @ " "git+https://github.com/Kitsunp/ml-cross-entropy.git@main" ) _EXTENDED_CCE_PARAMETERS = ( frozenset(inspect.signature(linear_cross_entropy).parameters) if linear_cross_entropy is not None else frozenset() ) # Do not mutate global TorchDynamo/torch.compile configuration here. The # compiler-safe CCE boundary remains optional and the caller owns any graph # policy it explicitly chooses; importing this modeling file must be harmless # when CCE or the optional Leviathan kernel is absent. def _require_extended_cce_options(*option_names: str) -> None: """Fail clearly only when an explicitly enabled CCE extension is unavailable.""" if linear_cross_entropy is None: return missing = [name for name in option_names if name not in _EXTENDED_CCE_PARAMETERS] if missing: raise ImportError( "The installed cut-cross-entropy package does not provide the enabled " f"extension options {missing}. Install the extended backend with " f"`pip install --force-reinstall --no-deps \"{_EXTENDED_CCE_INSTALL}\"`, " "or disable the corresponding NeoLLM flags." ) def compute_cce_loss( hidden_states, labels, lm_head_weight, lm_head_bias=None, pad_token_id=None, cce_impl="cce_kahan_full_c", use_mile_loss=False, mile_loss_gamma=1.0, mile_group_mask=None, use_mu_loss=False, mu_loss_lambda=1e-4, return_loss_metrics=False, ): """CCE loss with a compiler-visible custom-op boundary in the extended backend. The backend keeps data-dependent label compaction and Triton autotuning opaque while exposing one forward/backward operator to Dynamo. Older CCE installations retain their normal eager fallback and fail clearly only when an explicitly requested extension is unavailable. """ if linear_cross_entropy is None: raise ImportError( "NeoLLM was configured with `ntp_loss_backend='cce'`, but " "`cut_cross_entropy` is not installed. Install `cut-cross-entropy` " "or set `use_liger_kernel=True` / `ntp_loss_backend='liger'`." ) processed_labels = _prepare_lm_labels( labels, device=hidden_states.device, pad_token_id=pad_token_id ) extension_kwargs = {} if use_mile_loss: _require_extended_cce_options("mile_enabled", "mile_gamma") extension_kwargs.update( mile_enabled=True, mile_gamma=float(mile_loss_gamma), ) if mile_group_mask is not None: _require_extended_cce_options("mile_group_mask") extension_kwargs["mile_group_mask"] = mile_group_mask if use_mu_loss: _require_extended_cce_options("mu_loss_enabled", "mu_loss_lambda") extension_kwargs.update( mu_loss_enabled=True, mu_loss_lambda=float(mu_loss_lambda), ) if return_loss_metrics: _require_extended_cce_options("return_loss_metrics") extension_kwargs["return_loss_metrics"] = True return linear_cross_entropy( hidden_states, lm_head_weight, processed_labels, bias=lm_head_bias, shift=1, impl=cce_impl, reduction="mean", **extension_kwargs, ) def compute_liger_loss( hidden_states, labels, lm_head_weight, lm_head_bias=None, pad_token_id=None, liger_loss=None, ): """Liger FLCE NTP loss. Mirrors CCE shift=1 without materializing logits.""" if liger_loss is None: raise RuntimeError( "NeoLLM was configured with `ntp_loss_backend='liger'`, but the " "Liger loss module was not initialized." ) processed_labels = _prepare_lm_labels( labels, device=hidden_states.device, pad_token_id=pad_token_id ) # Match CCE shift=1: h_t predicts label_{t+1}. shift_hidden_states = hidden_states[..., :-1, :].contiguous() shift_labels = processed_labels[..., 1:].contiguous() hidden_size = shift_hidden_states.shape[-1] return liger_loss( lm_head_weight, shift_hidden_states.view(-1, hidden_size), shift_labels.view(-1), bias=lm_head_bias, ) def compute_tweo_loss( block_activations: Tuple[torch.Tensor, ...], tau: float, power: float, eps: float, attention_mask: Optional[torch.Tensor] = None, ) -> torch.Tensor: """ Transformers Without Extreme Outliers activation regularizer. Liang et al. (2026) apply a scaled Lp penalty to the final output activations of every Transformer block. The penalty is intentionally smooth rather than hinge-thresholded, so normal activations remain cheap while extreme outliers become expensive. Padding tokens are excluded when an attention mask is available; this keeps the regularizer aligned with the causal LM and auxiliary losses without changing tensor shapes. """ if not block_activations: raise ValueError("TWEO requires at least one Transformer block activation.") mask_weights = None valid_count = None if attention_mask is not None: mask_weights = attention_mask.to( device=block_activations[0].device, dtype=torch.float32 ) valid_count = mask_weights.sum().clamp_min(1.0) scale = float(tau) + float(eps) total = block_activations[0].new_zeros((), dtype=torch.float32) for activation in block_activations: penalty = (activation.float().abs() / scale).pow(float(power)) if mask_weights is None: total = total + penalty.mean() else: token_penalty = penalty.mean(dim=-1) total = total + (token_penalty * mask_weights).sum() / valid_count return total / len(block_activations) class NITPProjector(nn.Module): """ SwiGLU projection head for Next Implicit Token Prediction. Zhang et al. (2026) align projected final hidden states with shallow-layer implicit tokens. The projector is a training-only auxiliary head: gate: d -> 4d, up: d -> 4d, down: 4d -> d. """ def __init__(self, config: NeoLLMConfig): super().__init__() hidden_size = config.hidden_size intermediate_size = int(config.nitp_projector_intermediate_size) self.gate_proj = nn.Linear(hidden_size, intermediate_size, bias=False) self.up_proj = nn.Linear(hidden_size, intermediate_size, bias=False) self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: return self.down_proj( F.silu(self.gate_proj(hidden_states)) * self.up_proj(hidden_states) ) def compute_nitp_loss( final_hidden_states: torch.Tensor, target_hidden_states: torch.Tensor, labels: Optional[torch.LongTensor], attention_mask: Optional[torch.Tensor], projector: NITPProjector, pad_token_id: Optional[int] = None, ) -> torch.Tensor: """ Next Implicit Token Prediction loss. Predicts the shallow-layer representation of token t+1 from the final hidden state at token t. The target is detached (stop-gradient), the last predictive position is dropped, and padding / ignored labels are excluded. """ predicted = projector(final_hidden_states[:, :-1, :]) target = target_hidden_states[:, 1:, :].detach() per_token_loss = 1.0 - F.cosine_similarity( predicted.float(), target.float(), dim=-1, eps=1e-8, ) valid_mask = torch.ones_like(per_token_loss, dtype=torch.bool) if attention_mask is not None: current_valid = attention_mask[:, :-1].to(device=per_token_loss.device).bool() target_valid = attention_mask[:, 1:].to(device=per_token_loss.device).bool() valid_mask = valid_mask & current_valid & target_valid if labels is not None: shifted_labels = labels[:, 1:].to(device=per_token_loss.device) valid_mask = valid_mask & (shifted_labels != -100) if pad_token_id is not None: valid_mask = valid_mask & (shifted_labels != pad_token_id) valid_weights = valid_mask.to(dtype=per_token_loss.dtype) valid_count = valid_weights.sum().clamp_min(1.0) return (per_token_loss * valid_weights).sum() / valid_count class NITPTemporalTransition(nn.Module): """ Causal NITP-conditioned multiplicative transition. This module remains independent from ``NextLatDynamicsModel``. The current final hidden state proposes a latent update, while the detached NITP prediction of the next shallow state modulates that proposal: h_bar = RMSNorm_h(h_t) z_bar = RMSNorm_z(stopgrad(P_NITP(h_t))) u_h = SiLU(W_gate h_bar) * (W_up h_bar) b_tilde_k = b_k - mean_j(b_j) s_{z,k} = tanh(W_condition z_bar + b_tilde_k) g_k = 2 * sigmoid(a_k) delta_{t,k} = g_k * W_down(u_h * s_{z,k}) h_hat_{t+1} = h_t + delta_{t,k} The centered signed gate satisfies ``s_{z,k} in (-1, 1)``. Unlike a positive sigmoid gate, it can attenuate, activate, or reverse an individual contextual feature. A zero-initialized learned vector ``b_k`` modulates only the recurrent phase. Before indexing a horizon, the table is reparameterized as b_tilde_k = b_k - mean_j(b_j), so ``sum_k b_tilde_k = 0``. All horizons still share the same transition weights, but the bias table can encode only relative differences between rollout phases; a single offset shared by every horizon is removed exactly. A second zero-initialized scalar logit ``a_k`` supplies a positive bounded residual gain ``g_k = 2 sigmoid(a_k)`` for each horizon. Thus ``b_tilde_k`` controls the relative phase/channel pattern, while ``g_k`` controls only its global amplitude. At initialization ``b_k=0`` and ``a_k=0`` for every horizon, hence ``b_tilde_k=0``, ``g_k=1``, and the model starts exactly from the previous transition. The step index is deterministic and contains no future-token information. At initialization, ``W_condition=0`` implies ``s_{z,k}=0`` and therefore ``T_k(h, z)=h``. After training, a non-zero relative phase bias can provide a horizon-specific gate even when the condition projection is weak, but it cannot create the same bias offset at every horizon. The gain can rescale the resulting residual, yet neither phase nor gain creates an additive route through ``down_proj``: every output update remains multiplicatively tied to contextual features proposed by the current causal state. All projections are bias-free. Therefore the NITP condition has no additive route to the output and the transition still satisfies: F(0, z) = 0 dF(0, z) / dz = 0 The condition cannot construct a deep state by itself; it only controls features proposed by the current causal state. The module never materializes vocabulary logits. Parameter count, excluding the already-existing NITP projector: 4 * d * m + 2d + W * m + W where ``m = nitp_temporal_intermediate_size``, ``W`` is the temporal horizon, and ``2d`` comes from the two affine RMSNorm scales. With d=512, m=1024, and W=4 this is 2,102,276 trainable parameters: 4,100 more than the shared transition, of which only four are residual-gain parameters. """ def __init__(self, config: NeoLLMConfig): super().__init__() hidden_size = int(config.hidden_size) intermediate_size = int(config.nitp_temporal_intermediate_size) horizon = int(config.nitp_temporal_horizon) if horizon < 1: raise ValueError("nitp_temporal_horizon must be at least 1") # The two inputs have different semantics and are normalized # independently: a final-layer causal state and a predicted shallow # next-state condition produced by the NITP projector. self.context_norm = nn.RMSNorm(hidden_size, eps=config.rms_norm_eps) self.condition_norm = nn.RMSNorm(hidden_size, eps=config.rms_norm_eps) # Context path: proposes the features of the deep-state update. self.context_gate_proj = nn.Linear( hidden_size, intermediate_size, bias=False ) self.context_up_proj = nn.Linear( hidden_size, intermediate_size, bias=False ) # NITP-condition path: controls contextual features but has no additive # route to the output. The signed tanh gate is centered at zero; the # condition and the zero-mean relative phase bias modulate features # proposed by the current causal state. self.condition_modulation_proj = nn.Linear( hidden_size, intermediate_size, bias=False ) self.down_proj = nn.Linear(intermediate_size, hidden_size, bias=False) # Minimal phase conditioning. Each horizon receives one learned raw # bias in the shared intermediate gate. Forward reparameterizes this # table to zero mean over horizons, so only relative phase differences # are expressible through this branch. Zero initialization preserves # the exact previous transition at step 0 and avoids disturbing warmup. self.temporal_step_bias = nn.Parameter( torch.zeros(horizon, intermediate_size) ) # Minimal amplitude conditioning. One scalar logit per horizon is # mapped to a positive bounded gain in (0, 2). Zero initialization # gives exactly unit gain: 2 * sigmoid(0) = 1. self.temporal_step_gain_logits = nn.Parameter(torch.zeros(horizon)) def forward( self, current_states: torch.Tensor, next_shallow_conditions: torch.Tensor, step_index: int, ) -> torch.Tensor: context = self.context_norm(current_states) condition = self.condition_norm(next_shallow_conditions) context_update = ( F.silu(self.context_gate_proj(context)) * self.context_up_proj(context) ) # Remove the horizon-common bias mode without materializing a full # broadcast tensor. For horizon=1 this term is exactly zero, as there # is no relative phase to represent. relative_step_bias = ( self.temporal_step_bias[step_index] - self.temporal_step_bias.mean(dim=0) ) signed_modulation = torch.tanh( self.condition_modulation_proj(condition) + relative_step_bias ) delta = self.down_proj(context_update * signed_modulation) step_gain = 2.0 * torch.sigmoid( self.temporal_step_gain_logits[step_index] ) return current_states + step_gain * delta class NextLatDynamicsModel(nn.Module): """ Latent dynamics model p_psi(h_t, x_{t+1}) from NextLat. Teoh et al. (2026) use a simple MLP over the concatenation of the current final-layer hidden state and the teacher-forced next-token embedding. The residual output predicts the next final-layer hidden state. """ def __init__(self, config: NeoLLMConfig): super().__init__() hidden_size = int(config.hidden_size) input_dim = 2 * hidden_size hidden_dim = int(float(config.nextlat_dynamics_proj_factor) * input_dim) hidden_dim = max(128, 128 * round(hidden_dim / 128)) self.hidden_state_dropout = ( nn.Dropout(config.attention_dropout) if float(getattr(config, "attention_dropout", 0.0)) > 0.0 else nn.Identity() ) self.norm_x = nn.RMSNorm(input_dim, eps=config.rms_norm_eps) self.mlp = nn.Sequential( nn.Linear(input_dim, hidden_dim, bias=config.attention_bias), nn.GELU(), nn.Linear(hidden_dim, hidden_dim, bias=config.attention_bias), nn.GELU(), nn.Linear(hidden_dim, hidden_size, bias=config.attention_bias), ) def forward( self, current_states: torch.Tensor, next_token_embeds: torch.Tensor, ) -> torch.Tensor: hidden_states = self.hidden_state_dropout(current_states) x = torch.cat([next_token_embeds, hidden_states], dim=-1) return current_states + self.mlp(self.norm_x(x)) def _masked_mean(values: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: weights = mask.to(device=values.device, dtype=values.dtype) return (values * weights).sum() / weights.sum().clamp_min(1.0) def _masked_smooth_l1( predicted: torch.Tensor, target: torch.Tensor, mask: torch.Tensor, ) -> torch.Tensor: loss = F.smooth_l1_loss(predicted, target.detach(), reduction="none") weights = mask.to(device=loss.device, dtype=loss.dtype).unsqueeze(-1) return (loss * weights).sum() / weights.expand_as(loss).sum().clamp_min(1.0) def _valid_temporal_span_mask( valid_tokens: torch.Tensor, labels: torch.Tensor, span_length: int, eos_token_id: Optional[int], ) -> torch.Tensor: """ Return a mask for contiguous temporal spans that stay inside one document. A span includes the source state and every real future target used to evaluate the causal rollout. Padding / ignored labels invalidate the full span. EOS is permitted only at the final position, so no transition is trained across a packed-document boundary. """ token_windows = valid_tokens.unfold(dimension=1, size=span_length, step=1) span_mask = token_windows.all(dim=-1) if eos_token_id is not None and span_length > 1: label_windows = labels.unfold(dimension=1, size=span_length, step=1) # The last token may itself be EOS; all earlier positions must remain # inside the same document. span_mask = span_mask & (label_windows[..., :-1] != eos_token_id).all(dim=-1) return span_mask def _project_nitp_with_frozen_weights( hidden_states: torch.Tensor, projector: NITPProjector, ) -> torch.Tensor: """ Evaluate the current NITP projector while freezing only its parameters. ``torch.no_grad()`` would also remove the derivative with respect to ``hidden_states``. The temporal identification objective instead needs: d P_{sg(theta_P)}(h) / d h != 0 d P_{sg(theta_P)}(h) / d theta_P = 0. The explicit bias-free SwiGLU computation below therefore uses detached projector weights. Temporal gradients can move a rollout state into a region where the already-defined NITP representation remains valid, but they cannot turn the projector into a private code for the transition. """ gate = F.linear( hidden_states, projector.gate_proj.weight.detach(), bias=None, ) up = F.linear( hidden_states, projector.up_proj.weight.detach(), bias=None, ) return F.linear( F.silu(gate) * up, projector.down_proj.weight.detach(), bias=None, ) def _build_relational_temporal_similarity( predicted_by_step: list[torch.Tensor], condition_by_step: list[torch.Tensor], target_hidden_states: torch.Tensor, shallow_target_states: torch.Tensor, initial_hidden_states: torch.Tensor, common_length: int, ) -> Tuple[ torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, ]: """ Build the local ``W x W`` temporal-identification matrices. The trainable identification score keeps two conceptual halves: S_{k,j} = 0.5 * [Gamma_{k,j} + Z_{k,j}], where the geometric half is itself an equal-norm concatenation: Gamma_{k,j} = 0.5 * [A_{k,j} + R_{k,j}]. The expanded score is therefore S_{k,j} = 0.25 A_{k,j} + 0.25 R_{k,j} + 0.5 Z_{k,j}. ``A`` is the direct accumulated-displacement cosine: u_hat_k = normalize(h_hat_k - h_0) u_j^* = normalize(stopgrad(h_j^* - h_0)) A_{k,j} = u_hat_k^T u_j^*. ``R`` compares relational angular profiles instead of forcing the rollout to copy the teacher's local tangent. First compute G^*_{j,r} = (u_j^*)^T u_r^*, rho_hat_k = normalize(A_{k,:}), rho_j^* = normalize(G^*_{j,:}), R_{k,j} = rho_hat_k^T rho_j^*. If all rollout displacements collapse to one direction while the target trajectory has distinct directions, the rows ``A_{k,:}`` become nearly identical and cannot match the distinct target profiles ``G^*_{j,:}``. Conversely, if the target trajectory is genuinely parallel, its profiles are also similar, so this term does not invent artificial diversity. ``Z`` is the frozen-weight NITP shallow-signature cosine. Projector parameters remain stop-gradient while gradients through the rollout state are preserved by ``_project_nitp_with_frozen_weights``. No tunable mixing coefficient is introduced: 1/2 and 1/4 follow exactly from concatenating independently normalized blocks. The implementation deliberately evaluates the W-by-W cosines row-by-row instead of creating a ``[B,T,W,D]`` FP32 tensor, keeping peak memory small for W=4. """ horizon = len(predicted_by_step) initial_common = initial_hidden_states[:, :common_length, :] detached_initial = initial_common.detach() displacement_rows = [] for predicted_states in predicted_by_step: predicted_displacement = ( predicted_states[:, :common_length, :] - initial_common ) row = [] for offset in range(1, horizon + 1): target_displacement = ( target_hidden_states[ :, offset : offset + common_length, : ].detach() - detached_initial ) row.append( F.cosine_similarity( predicted_displacement.float(), target_displacement.float(), dim=-1, eps=1.0e-8, ) ) displacement_rows.append(torch.stack(row, dim=-1)) displacement_similarity = torch.stack(displacement_rows, dim=-2) target_gram_rows = [] for left_offset in range(1, horizon + 1): left_displacement = ( target_hidden_states[ :, left_offset : left_offset + common_length, : ].detach() - detached_initial ) row = [] for right_offset in range(1, horizon + 1): right_displacement = ( target_hidden_states[ :, right_offset : right_offset + common_length, : ].detach() - detached_initial ) row.append( F.cosine_similarity( left_displacement.float(), right_displacement.float(), dim=-1, eps=1.0e-8, ) ) target_gram_rows.append(torch.stack(row, dim=-1)) target_gram = torch.stack(target_gram_rows, dim=-2) predicted_profiles = F.normalize( displacement_similarity, p=2.0, dim=-1, eps=1.0e-8, ) target_profiles = F.normalize( target_gram, p=2.0, dim=-1, eps=1.0e-8, ) relational_similarity = torch.einsum( "btkr,btjr->btkj", predicted_profiles, target_profiles, ) nitp_rows = [] for predicted_condition in condition_by_step: predicted_condition = predicted_condition[:, :common_length, :] row = [] for offset in range(1, horizon + 1): target_condition = shallow_target_states[ :, offset : offset + common_length, : ].detach() row.append( F.cosine_similarity( predicted_condition.float(), target_condition.float(), dim=-1, eps=1.0e-8, ) ) nitp_rows.append(torch.stack(row, dim=-1)) nitp_signature_similarity = torch.stack(nitp_rows, dim=-2) geometric_similarity = 0.5 * ( displacement_similarity + relational_similarity ) joint_similarity = 0.5 * ( geometric_similarity + nitp_signature_similarity ) return ( joint_similarity, displacement_similarity, relational_similarity, nitp_signature_similarity, ) def _temporal_matrix_statistics( matrix: torch.Tensor, window_mask: torch.Tensor, ) -> torch.Tensor: """ Return ``[diagonal, off_diagonal, margin, positive_margin_fraction]``. The margin is the diagonal score minus the hardest incorrect candidate in each row. ``positive_margin_fraction`` prevents a positive mean margin from hiding a large fraction of individually misidentified rows. """ horizon = int(matrix.shape[-1]) zero = matrix.new_zeros((), dtype=torch.float32) diagonal = matrix.diagonal(dim1=-2, dim2=-1) diagonal_mean = _masked_mean(diagonal.mean(dim=-1), window_mask) if horizon <= 1: return torch.stack((diagonal_mean, zero, zero, zero)) identity = torch.eye(horizon, dtype=torch.bool, device=matrix.device) identity = identity.view(1, 1, horizon, horizon) off_diagonal = matrix.masked_select(~identity).view( matrix.shape[0], matrix.shape[1], horizon * (horizon - 1), ) off_diagonal_mean = _masked_mean( off_diagonal.mean(dim=-1), window_mask, ) hard_negative = matrix.masked_fill( identity, torch.finfo(matrix.dtype).min, ).max(dim=-1).values row_margin = diagonal - hard_negative margin = _masked_mean(row_margin.mean(dim=-1), window_mask) positive_fraction = _masked_mean( (row_margin > 0.0).float().mean(dim=-1), window_mask, ) return torch.stack( (diagonal_mean, off_diagonal_mean, margin, positive_fraction) ) def _temporal_neighbor_cosines( matrix: torch.Tensor, window_mask: torch.Tensor, ) -> torch.Tensor: """Average matrix entries grouped by temporal distance ``|k-j|``.""" horizon = int(matrix.shape[-1]) if horizon <= 1: return matrix.new_zeros((0,), dtype=torch.float32) indices = torch.arange(horizon, device=matrix.device) temporal_distance = (indices[:, None] - indices[None, :]).abs() values = [] for offset in range(1, horizon): offset_mask = temporal_distance.eq(offset).view( 1, 1, horizon, horizon ) selected = matrix.masked_select(offset_mask).view( matrix.shape[0], matrix.shape[1], -1, ) values.append(_masked_mean(selected.mean(dim=-1), window_mask)) return torch.stack(values) def _pairwise_cosine_matrix(states: torch.Tensor) -> torch.Tensor: """Return the within-trajectory cosine Gram matrix for ``[B,T,W,D]``.""" unit = F.normalize(states.float(), p=2.0, dim=-1, eps=1.0e-8) return torch.einsum("btkd,btjd->btkj", unit, unit) def _pairwise_cosine_from_steps( states_by_step: list[torch.Tensor], origin: Optional[torch.Tensor] = None, ) -> torch.Tensor: """Low-peak-memory pair matrix from W tensors shaped ``[B,T,D]``.""" rows = [] for left in states_by_step: left_value = left if origin is None else left - origin row = [] for right in states_by_step: right_value = right if origin is None else right - origin row.append( F.cosine_similarity( left_value.float(), right_value.float(), dim=-1, eps=1.0e-8, ) ) rows.append(torch.stack(row, dim=-1)) return torch.stack(rows, dim=-2) def _off_diagonal_pair_mean( matrix: torch.Tensor, window_mask: torch.Tensor, ) -> torch.Tensor: horizon = int(matrix.shape[-1]) zero = matrix.new_zeros((), dtype=torch.float32) if horizon <= 1: return zero identity = torch.eye(horizon, dtype=torch.bool, device=matrix.device) selected = matrix.masked_select( ~identity.view(1, 1, horizon, horizon) ).view( matrix.shape[0], matrix.shape[1], horizon * (horizon - 1), ) return _masked_mean(selected.mean(dim=-1), window_mask) def _temporal_two_token_readout( student_states: torch.Tensor, teacher_states: torch.Tensor, candidate_labels: torch.LongTensor, candidate_valid: torch.Tensor, target_index: int, base_mask: torch.Tensor, lm_head_weight: torch.Tensor, lm_head_bias: Optional[torch.Tensor], ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]: """ Cheap functional anchor over the local temporal candidate set. For the real state at horizon ``k`` the candidate set contains only the tokens attached to the W local future horizons. The strongest *different* temporal token under the real state is selected as the competitor. The rollout state then matches the teacher's binary distribution over [token at the correct horizon, strongest incorrect temporal token]. This avoids a [B,T,W,V] or [B,T,V] vocabulary projection. It evaluates W teacher dot-products and two student dot-products, using detached LM-head rows so the auxiliary loss trains the temporal trajectory rather than the vocabulary head. No new temperature, loss weight, or trainable parameter is introduced. """ horizon = int(candidate_labels.shape[-1]) zero = student_states.new_zeros((), dtype=torch.float32) if horizon <= 1 or int(student_states.shape[1]) == 0: return zero, zero, zero, zero detached_weight = lm_head_weight.detach() detached_bias = lm_head_bias.detach() if lm_head_bias is not None else None teacher = teacher_states.detach() teacher_scores_by_candidate = [] for candidate_index in range(horizon): token_ids = candidate_labels[..., candidate_index].clamp_min(0) token_rows = F.embedding(token_ids, detached_weight) score = (teacher * token_rows).sum(dim=-1, dtype=torch.float32) if detached_bias is not None: score = score + detached_bias[token_ids].float() teacher_scores_by_candidate.append(score) teacher_scores = torch.stack(teacher_scores_by_candidate, dim=-1) target_ids = candidate_labels[..., target_index] safe_target_ids = target_ids.clamp_min(0) target_valid = candidate_valid[..., target_index] target_teacher_score = teacher_scores[..., target_index] candidate_indices = torch.arange( horizon, device=candidate_labels.device, ).view(1, 1, horizon) competitor_mask = ( candidate_valid & candidate_indices.ne(int(target_index)) & candidate_labels.ne(target_ids.unsqueeze(-1)) ) has_competitor = competitor_mask.any(dim=-1) masked_teacher_scores = teacher_scores.masked_fill( ~competitor_mask, torch.finfo(teacher_scores.dtype).min, ) competitor_index = masked_teacher_scores.argmax(dim=-1) competitor_ids = candidate_labels.gather( dim=-1, index=competitor_index.unsqueeze(-1), ).squeeze(-1) safe_competitor_ids = competitor_ids.clamp_min(0) competitor_teacher_score = teacher_scores.gather( dim=-1, index=competitor_index.unsqueeze(-1), ).squeeze(-1) # Compute the two student logits sequentially so only one gathered # [B,T,H] LM-head block is live at a time. target_rows = F.embedding(safe_target_ids, detached_weight) target_student_score = ( student_states * target_rows ).sum(dim=-1, dtype=torch.float32) competitor_rows = F.embedding(safe_competitor_ids, detached_weight) competitor_student_score = ( student_states * competitor_rows ).sum(dim=-1, dtype=torch.float32) if detached_bias is not None: target_student_score = ( target_student_score + detached_bias[safe_target_ids].float() ) competitor_student_score = ( competitor_student_score + detached_bias[safe_competitor_ids].float() ) teacher_pair_logits = torch.stack( (target_teacher_score, competitor_teacher_score), dim=-1, ) student_pair_logits = torch.stack( (target_student_score, competitor_student_score), dim=-1, ) readout_mask = base_mask & target_valid & has_competitor teacher_log_probs = F.log_softmax(teacher_pair_logits.detach(), dim=-1) student_log_probs = F.log_softmax(student_pair_logits, dim=-1) per_token_kl = F.kl_div( student_log_probs, teacher_log_probs, log_target=True, reduction="none", ).sum(dim=-1) readout_loss = _masked_mean(per_token_kl, readout_mask) with torch.no_grad(): readout_agreement = _masked_mean( ( teacher_pair_logits.argmax(dim=-1) == student_pair_logits.argmax(dim=-1) ).float(), readout_mask, ) teacher_margin = ( target_teacher_score - competitor_teacher_score ) student_margin = ( target_student_score - competitor_student_score ) readout_margin_error = _masked_mean( (student_margin - teacher_margin).abs(), readout_mask, ) readout_valid_fraction = readout_mask.float().mean() return ( readout_loss, readout_agreement, readout_margin_error, readout_valid_fraction, ) NITP_TEMPORAL_STEP_METRIC_NAMES = ( # Rows 0..12 are stable for backward-compatible trainer logging. "state_loss", "composite_step_loss", "delta_cosine", "delta_norm_ratio", "rollout_nitp_cosine", "readout_loss", "readout_agreement", "readout_margin_error", "readout_valid_fraction", "cross_example_pred_displacement_cosine", "cross_example_target_displacement_cosine", "predicted_delta_turn_cosine", "target_delta_turn_cosine", # Append-only component diagnostics for the new cheap moment objective. "raw_step_loss", "mean_error_loss", "std_match_loss", ) def compute_nitp_temporal_objective( hidden_states: torch.Tensor, shallow_target_states: torch.Tensor, labels: torch.LongTensor, attention_mask: Optional[torch.Tensor], transition_model: NITPTemporalTransition, nitp_projector: NITPProjector, lm_head_weight: torch.Tensor, lm_head_bias: Optional[torch.Tensor], horizon: int, dynamics_weight: float, identification_weight: float, similarity_temperature: float, target_temperature: float, pad_token_id: Optional[int] = None, eos_token_id: Optional[int] = None, ) -> Tuple[torch.Tensor, ...]: """ Causal Temporal-NITP with a cheap first-order trajectory objective. The autonomous rollout and relational identification are unchanged: h_hat_0 = h_0, z_hat_k = P_{sg(theta_P)}(h_hat_{k-1}), h_hat_k = T_k(h_hat_{k-1}, sg(z_hat_k)), S = 0.25 A + 0.25 R + 0.5 Z. Only the *dynamics objective* changes. The old objective supervised final states alone. The new objective is the fixed arithmetic mean L_dyn = (L_position + L_step + L_readout) / 3. ``L_position`` is the original Smooth-L1 state alignment. ``L_step`` uses target-RMS-normalized local increments and is now L_step = L_raw + L_mean_error + 0.5 * L_std. ``L_raw`` keeps the full per-window correspondence. ``L_mean_error`` applies Smooth-L1 only to the valid-window mean prediction error, directly penalizing a horizon-common direction that is absent from the targets. ``L_std`` matches the per-channel standard deviation around those means, retaining the useful lesson from the removed explicit centered loss: the rollout must preserve context-dependent dispersion instead of shrinking every example toward its batch mean. Means and second moments are reductions to [1,1,H]; no centered [B,S,H] tensors or second full-size Smooth-L1 branch are materialized. ``L_readout`` matches a two-token temporal readout selected from the existing W future labels; it protects functionally sensitive LM-head directions without projecting to the full vocabulary. The constants 1/3 and 0.5 are not configurable hyperparameters. The first defines an equal mean of the three trajectory views; the second keeps the dispersion match auxiliary to the exact per-window and common-error terms. Existing ``dynamics_weight`` and ``identification_weight`` remain the only outer weights. Packed diagnostics are intentionally compact and non-redundant: ``component_metrics`` [3,4] rows: displacement A, relational R, NITP signature Z. cols: diagonal, off-diagonal, hardest-negative margin, positive-margin fraction. The joint diagonal/off-diagonal are omitted because they are exactly reconstructed as 0.25*A + 0.25*R + 0.5*Z; joint positive-margin fraction is identical to joint top-1. ``step_metrics`` [16,W] Rows 0..12 preserve the previous order for logger compatibility: state loss, composite normalized step loss, delta cosine, delta norm ratio, rollout-NITP cosine, two-token readout KL, readout rank agreement, readout margin error, readout valid fraction, cross-example predicted displacement cosine, cross-example target displacement cosine, predicted consecutive-delta cosine, target consecutive-delta cosine. Rows 13..15 append the three step-loss components: raw local loss, common mean-error loss, and standard-deviation matching loss. The composite row is exactly raw + mean_error + 0.5*std. Aggregate component means are reconstructible by averaging the appended rows over W. ``offset_metrics`` [3,W-1] neighbor similarities for A, R, and Z. The joint neighbor profile is omitted because it is exactly reconstructed from these three rows. ``scalar_metrics`` [5] path-length ratio, predicted and target absolute-state pair cosine, predicted and target centered-displacement pair cosine. Pair gaps are omitted because they are simple differences. The previous counterfactual trajectory-specificity rollout is removed. It required an additional W-step transition rollout and is replaced by the much cheaper per-step readout agreement plus cross-example diversity. """ horizon = int(horizon) seq_len = int(hidden_states.shape[1]) max_horizon = max(min(horizon, seq_len - 1), 0) zero = hidden_states.new_zeros((), dtype=torch.float32) if max_horizon == 0: return ( *((zero,) * 9), zero.new_zeros((3, 4)), zero.new_zeros((16, 0)), zero.new_zeros((3, 0)), zero.new_zeros((5,)), ) device = hidden_states.device labels_on_device = labels.to(device=device) if attention_mask is None: valid_tokens = torch.ones_like(labels_on_device, dtype=torch.bool) else: valid_tokens = attention_mask.to(device=device).bool() valid_tokens = valid_tokens & (labels_on_device != -100) if pad_token_id is not None: valid_tokens = valid_tokens & (labels_on_device != pad_token_id) # The functional readout needs one token beyond the farthest hidden-state # target because h_{t+k} predicts label_{t+k+1}. Candidate labels are the W # true tokens attached to the local temporal horizons, so selecting the # strongest incorrect one costs O(W*H), not O(V*H). functional_length = max(seq_len - max_horizon - 1, 0) if functional_length > 0: candidate_labels = torch.stack( [ labels_on_device[ :, offset + 1 : offset + 1 + functional_length ] for offset in range(1, max_horizon + 1) ], dim=-1, ) candidate_valid = torch.stack( [ valid_tokens[ :, offset + 1 : offset + 1 + functional_length ] for offset in range(1, max_horizon + 1) ], dim=-1, ) functional_window_mask = _valid_temporal_span_mask( valid_tokens=valid_tokens, labels=labels_on_device, span_length=max_horizon + 2, eos_token_id=eos_token_id, ) else: candidate_labels = None candidate_valid = None functional_window_mask = None current_states = hidden_states shifted_targets = hidden_states predicted_by_step: list[torch.Tensor] = [] condition_by_step: list[torch.Tensor] = [] step_masks: list[torch.Tensor] = [] state_loss_by_step: list[torch.Tensor] = [] normalized_step_loss_by_step: list[torch.Tensor] = [] raw_step_loss_by_step: list[torch.Tensor] = [] mean_error_loss_by_step: list[torch.Tensor] = [] std_match_loss_by_step: list[torch.Tensor] = [] delta_cosine_by_step: list[torch.Tensor] = [] delta_norm_ratio_by_step: list[torch.Tensor] = [] rollout_nitp_cosine_by_step: list[torch.Tensor] = [] readout_loss_by_step: list[torch.Tensor] = [] readout_agreement_by_step: list[torch.Tensor] = [] readout_margin_error_by_step: list[torch.Tensor] = [] readout_valid_fraction_by_step: list[torch.Tensor] = [] predicted_turn_by_step: list[torch.Tensor] = [] target_turn_by_step: list[torch.Tensor] = [] state_loss_total = zero normalized_step_loss_total = zero readout_loss_total = zero path_predicted_norm_sum = zero path_target_norm_sum = zero previous_predicted_delta: Optional[torch.Tensor] = None previous_real_delta: Optional[torch.Tensor] = None for step in range(max_horizon): source_states = current_states[:, :-1, :] shifted_targets = shifted_targets[:, 1:, :] next_shallow_signature = _project_nitp_with_frozen_weights( source_states, nitp_projector, ) predicted_states = transition_model( source_states, next_shallow_signature.detach(), step_index=step, ) step_mask = _valid_temporal_span_mask( valid_tokens=valid_tokens, labels=labels_on_device, span_length=step + 2, eos_token_id=eos_token_id, ) step_masks.append(step_mask) step_state_loss = _masked_smooth_l1( predicted=predicted_states, target=shifted_targets, mask=step_mask, ) state_loss_total = state_loss_total + step_state_loss state_loss_by_step.append(step_state_loss.detach().float()) real_source_states = hidden_states[ :, step : step + predicted_states.shape[1], : ].detach() predicted_delta = predicted_states - source_states real_delta = shifted_targets.detach() - real_source_states # Scale each target delta to unit RMS. This is data-derived and detached, # so there is no new tuned scale and no radial shortcut. target_delta_rms = real_delta.float().pow(2).mean( dim=-1, keepdim=True, ).sqrt().clamp_min(1.0e-4) normalized_predicted_delta = predicted_delta.float() / target_delta_rms normalized_real_delta = real_delta.float() / target_delta_rms # Absolute view: preserve the legitimate common temporal component as # well as each local direction and magnitude. raw_step_error = F.smooth_l1_loss( normalized_predicted_delta, normalized_real_delta, reduction="none", ).mean(dim=-1) raw_step_loss = _masked_mean(raw_step_error, step_mask) # Cheap moment supervision over valid windows. Unlike the removed # explicit centered branch, this directly penalizes the common *error* # while preserving the full raw loss. The standard-deviation term keeps # the context-specific spread from collapsing. Only [1,1,H] reductions # survive; no centered [B,S,H] tensors are constructed. step_weights = step_mask.to( device=normalized_predicted_delta.device, dtype=normalized_predicted_delta.dtype, ).unsqueeze(-1) valid_step_count = step_weights.sum( dim=(0, 1), keepdim=True, ).clamp_min(1.0) predicted_step_mean = ( normalized_predicted_delta * step_weights ).sum(dim=(0, 1), keepdim=True) / valid_step_count target_step_mean = ( normalized_real_delta * step_weights ).sum(dim=(0, 1), keepdim=True) / valid_step_count mean_error_loss = F.smooth_l1_loss( predicted_step_mean, target_step_mean.detach(), reduction="mean", ) predicted_second_moment = ( normalized_predicted_delta.square() * step_weights ).sum(dim=(0, 1), keepdim=True) / valid_step_count target_second_moment = ( normalized_real_delta.square() * step_weights ).sum(dim=(0, 1), keepdim=True) / valid_step_count # E[x^2] - E[x]^2 is clamped only for roundoff. The same epsilon on # both paths makes exactly constant prediction/target channels match. predicted_step_std = ( predicted_second_moment - predicted_step_mean.square() ).clamp_min(0.0).add(1.0e-6).sqrt() target_step_std = ( target_second_moment - target_step_mean.square() ).clamp_min(0.0).add(1.0e-6).sqrt() std_match_loss = F.smooth_l1_loss( predicted_step_std, target_step_std.detach(), reduction="mean", ) # Preserve the full local supervision, add direct pressure on the # horizon-common error, and retain a smaller dispersion constraint. step_delta_loss = ( raw_step_loss + mean_error_loss + 0.5 * std_match_loss ) normalized_step_loss_total = ( normalized_step_loss_total + step_delta_loss ) normalized_step_loss_by_step.append(step_delta_loss.detach().float()) raw_step_loss_by_step.append(raw_step_loss.detach().float()) mean_error_loss_by_step.append(mean_error_loss.detach().float()) std_match_loss_by_step.append(std_match_loss.detach().float()) if functional_length > 0: ( step_readout_loss, step_readout_agreement, step_readout_margin_error, step_readout_valid_fraction, ) = _temporal_two_token_readout( student_states=predicted_states[:, :functional_length, :], teacher_states=shifted_targets[:, :functional_length, :], candidate_labels=candidate_labels, candidate_valid=candidate_valid, target_index=step, base_mask=functional_window_mask, lm_head_weight=lm_head_weight, lm_head_bias=lm_head_bias, ) else: step_readout_loss = zero step_readout_agreement = zero step_readout_margin_error = zero step_readout_valid_fraction = zero readout_loss_total = readout_loss_total + step_readout_loss readout_loss_by_step.append(step_readout_loss.detach().float()) readout_agreement_by_step.append( step_readout_agreement.detach().float() ) readout_margin_error_by_step.append( step_readout_margin_error.detach().float() ) readout_valid_fraction_by_step.append( step_readout_valid_fraction.detach().float() ) with torch.no_grad(): step_delta_cosine = F.cosine_similarity( predicted_delta.float(), real_delta.float(), dim=-1, eps=1.0e-8, ) delta_cosine_by_step.append( _masked_mean(step_delta_cosine, step_mask) ) predicted_norm = predicted_delta.float().norm(dim=-1) real_norm = real_delta.float().norm(dim=-1) predicted_norm_mean = _masked_mean(predicted_norm, step_mask) real_norm_mean = _masked_mean(real_norm, step_mask) delta_norm_ratio_by_step.append( predicted_norm_mean / real_norm_mean.clamp_min(1.0e-8) ) step_weights = step_mask.to(dtype=torch.float32) path_predicted_norm_sum = path_predicted_norm_sum + ( predicted_norm * step_weights ).sum() path_target_norm_sum = path_target_norm_sum + ( real_norm * step_weights ).sum() shallow_target = shallow_target_states[:, step + 1 :, :] step_nitp_cosine = F.cosine_similarity( next_shallow_signature.detach().float(), shallow_target.float(), dim=-1, eps=1.0e-8, ) rollout_nitp_cosine_by_step.append( _masked_mean(step_nitp_cosine, step_mask) ) if previous_predicted_delta is None: predicted_turn_by_step.append(zero) target_turn_by_step.append(zero) else: previous_predicted_aligned = previous_predicted_delta[ :, : predicted_delta.shape[1], : ] previous_real_aligned = previous_real_delta[ :, : real_delta.shape[1], : ] predicted_turn = F.cosine_similarity( previous_predicted_aligned.float(), predicted_delta.float(), dim=-1, eps=1.0e-8, ) target_turn = F.cosine_similarity( previous_real_aligned.float(), real_delta.float(), dim=-1, eps=1.0e-8, ) predicted_turn_by_step.append( _masked_mean(predicted_turn, step_mask) ) target_turn_by_step.append( _masked_mean(target_turn, step_mask) ) previous_predicted_delta = predicted_delta.detach() previous_real_delta = real_delta.detach() predicted_by_step.append(predicted_states) condition_by_step.append(next_shallow_signature) current_states = predicted_states inv_horizon = 1.0 / float(max_horizon) state_loss = state_loss_total * inv_horizon normalized_step_loss = normalized_step_loss_total * inv_horizon readout_loss = readout_loss_total * inv_horizon trajectory_loss = ( state_loss + normalized_step_loss + readout_loss ) / 3.0 state_step_tensor = torch.stack(state_loss_by_step) normalized_step_loss_tensor = torch.stack( normalized_step_loss_by_step ) raw_step_loss_tensor = torch.stack(raw_step_loss_by_step) mean_error_loss_tensor = torch.stack(mean_error_loss_by_step) std_match_loss_tensor = torch.stack(std_match_loss_by_step) delta_step_tensor = torch.stack(delta_cosine_by_step) delta_norm_ratio_tensor = torch.stack(delta_norm_ratio_by_step) rollout_nitp_step_tensor = torch.stack(rollout_nitp_cosine_by_step) readout_loss_tensor = torch.stack(readout_loss_by_step) readout_agreement_tensor = torch.stack(readout_agreement_by_step) readout_margin_error_tensor = torch.stack( readout_margin_error_by_step ) readout_valid_fraction_tensor = torch.stack( readout_valid_fraction_by_step ) predicted_turn_tensor = torch.stack(predicted_turn_by_step) target_turn_tensor = torch.stack(target_turn_by_step) path_length_ratio = ( path_predicted_norm_sum / path_target_norm_sum.clamp_min(1.0e-8) ) common_length = seq_len - max_horizon if common_length <= 0: zero_steps = zero.new_zeros((max_horizon,)) step_metrics = torch.stack( ( state_step_tensor, normalized_step_loss_tensor, delta_step_tensor, delta_norm_ratio_tensor, rollout_nitp_step_tensor, readout_loss_tensor, readout_agreement_tensor, readout_margin_error_tensor, readout_valid_fraction_tensor, zero_steps, zero_steps, predicted_turn_tensor, target_turn_tensor, raw_step_loss_tensor, mean_error_loss_tensor, std_match_loss_tensor, ) ) weighted_total = float(dynamics_weight) * trajectory_loss return ( weighted_total, state_loss, normalized_step_loss, readout_loss, zero, zero, zero, zero, zero, zero.new_zeros((3, 4)), step_metrics, zero.new_zeros((3, max(max_horizon - 1, 0))), torch.stack((path_length_ratio, zero, zero, zero, zero)), ) ( similarity, displacement_similarity, relational_similarity, nitp_signature_similarity, ) = _build_relational_temporal_similarity( predicted_by_step=predicted_by_step, condition_by_step=condition_by_step, target_hidden_states=hidden_states, shallow_target_states=shallow_target_states, initial_hidden_states=hidden_states, common_length=common_length, ) window_mask = _valid_temporal_span_mask( valid_tokens=valid_tokens, labels=labels_on_device, span_length=max_horizon + 1, eos_token_id=eos_token_id, ) valid_window_fraction = window_mask.float().mean() offsets = torch.arange(max_horizon, device=device) distance = (offsets[:, None] - offsets[None, :]).abs().float() target_scores = -distance / float(target_temperature) row_target = F.softmax(target_scores, dim=-1) col_target = F.softmax(target_scores, dim=0) identification_logits = similarity / float(similarity_temperature) row_log_probs = F.log_softmax(identification_logits, dim=-1) col_log_probs = F.log_softmax(identification_logits, dim=-2) row_loss = -(row_target * row_log_probs).sum(dim=-1).mean(dim=-1) col_loss = -(col_target * col_log_probs).sum(dim=-2).mean(dim=-1) identification_loss = _masked_mean( 0.5 * (row_loss + col_loss), window_mask, ) row_prediction = similarity.argmax(dim=-1) expected_offsets = offsets.view(1, 1, -1) temporal_top1_accuracy = _masked_mean( (row_prediction == expected_offsets).float().mean(dim=-1), window_mask, ) mean_absolute_offset = _masked_mean( (row_prediction - expected_offsets).abs().float().mean(dim=-1), window_mask, ) identification_margin = _temporal_matrix_statistics( similarity.detach(), window_mask, )[2] # A/R/Z component rows only. Joint diagonal/off-diagonal and neighbor # profiles are exact fixed-weight reconstructions and are not logged. component_matrices = ( displacement_similarity.detach(), relational_similarity.detach(), nitp_signature_similarity.detach(), ) component_metrics = torch.stack( [ _temporal_matrix_statistics(component, window_mask) for component in component_matrices ] ) offset_metrics = torch.stack( [ _temporal_neighbor_cosines(component, window_mask) for component in component_matrices ] ) with torch.no_grad(): initial_common = hidden_states[:, :common_length, :].detach() predicted_state_steps = [ state[:, :common_length, :].detach() for state in predicted_by_step ] target_state_steps = [ hidden_states[ :, offset : offset + common_length, : ].detach() for offset in range(1, max_horizon + 1) ] rollout_pair_cosine = _off_diagonal_pair_mean( _pairwise_cosine_from_steps(predicted_state_steps), window_mask, ) target_pair_cosine = _off_diagonal_pair_mean( _pairwise_cosine_from_steps(target_state_steps), window_mask, ) rollout_displacement_pair_cosine = _off_diagonal_pair_mean( _pairwise_cosine_from_steps( predicted_state_steps, origin=initial_common, ), window_mask, ) target_displacement_pair_cosine = _off_diagonal_pair_mean( _pairwise_cosine_from_steps( target_state_steps, origin=initial_common, ), window_mask, ) cross_predicted_by_step = [] cross_target_by_step = [] if int(hidden_states.shape[0]) > 1: for step in range(max_horizon): step_mask_common = step_masks[step][:, :common_length] paired_mask = step_mask_common & torch.roll( step_mask_common, shifts=1, dims=0, ) predicted_displacement = ( predicted_by_step[step][:, :common_length, :].detach() - initial_common ) target_displacement = ( hidden_states[ :, step + 1 : step + 1 + common_length, : ].detach() - initial_common ) cross_predicted = F.cosine_similarity( predicted_displacement.float(), torch.roll( predicted_displacement.float(), shifts=1, dims=0, ), dim=-1, eps=1.0e-8, ) cross_target = F.cosine_similarity( target_displacement.float(), torch.roll( target_displacement.float(), shifts=1, dims=0, ), dim=-1, eps=1.0e-8, ) cross_predicted_by_step.append( _masked_mean(cross_predicted, paired_mask) ) cross_target_by_step.append( _masked_mean(cross_target, paired_mask) ) else: cross_predicted_by_step = [zero] * max_horizon cross_target_by_step = [zero] * max_horizon cross_predicted_tensor = torch.stack(cross_predicted_by_step) cross_target_tensor = torch.stack(cross_target_by_step) step_metrics = torch.stack( ( state_step_tensor, normalized_step_loss_tensor, delta_step_tensor, delta_norm_ratio_tensor, rollout_nitp_step_tensor, readout_loss_tensor, readout_agreement_tensor, readout_margin_error_tensor, readout_valid_fraction_tensor, cross_predicted_tensor, cross_target_tensor, predicted_turn_tensor, target_turn_tensor, raw_step_loss_tensor, mean_error_loss_tensor, std_match_loss_tensor, ) ) scalar_metrics = torch.stack( ( path_length_ratio, rollout_pair_cosine, target_pair_cosine, rollout_displacement_pair_cosine, target_displacement_pair_cosine, ) ) weighted_total = ( float(dynamics_weight) * trajectory_loss + float(identification_weight) * identification_loss ) return ( weighted_total, state_loss, normalized_step_loss, readout_loss, identification_loss, temporal_top1_accuracy, mean_absolute_offset, identification_margin, valid_window_fraction, component_metrics, step_metrics, offset_metrics, scalar_metrics, ) def _categorical_kl_loss( teacher_logits: torch.Tensor, student_logits: torch.Tensor, mask: torch.Tensor, ) -> torch.Tensor: log_teacher = F.log_softmax(teacher_logits.detach(), dim=-1) log_student = F.log_softmax(student_logits, dim=-1) per_token = F.kl_div( log_student, log_teacher, log_target=True, reduction="none", ).sum(dim=-1) return _masked_mean(per_token, mask) def compute_nextlat_loss( hidden_states: torch.Tensor, token_embeds: torch.Tensor, labels: torch.LongTensor, attention_mask: Optional[torch.Tensor], dynamics_model: NextLatDynamicsModel, lm_head_weight: torch.Tensor, lm_head_bias: Optional[torch.Tensor], horizon: int, mse_weight: float, kl_weight: float, pad_token_id: Optional[int] = None, ) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: """ NextLat training objective from Teoh et al. (2026). For each teacher-forced rollout step i <= d, p_psi predicts h_{t+i} recursively from h_t and x_{t+1:t+i}. Smooth L1 uses stop-gradient targets; KL compares the frozen LM-head distribution of sg[h_{t+i}] to the frozen LM-head distribution of the predicted state. """ horizon = int(horizon) seq_len = hidden_states.shape[1] max_horizon = max(min(horizon, seq_len - 1), 0) zero = hidden_states.new_zeros(()) if max_horizon == 0: return zero, zero, zero if attention_mask is None: valid_tokens = torch.ones( labels.shape, dtype=torch.bool, device=hidden_states.device ) else: valid_tokens = attention_mask.to(device=hidden_states.device).bool() label_tokens = labels.to(device=hidden_states.device) valid_tokens = valid_tokens & (label_tokens != -100) if pad_token_id is not None: valid_tokens = valid_tokens & (label_tokens != pad_token_id) pred_states = hidden_states next_token_embeds = token_embeds target_states = hidden_states mse_total = zero kl_total = zero frozen_weight = lm_head_weight.detach() frozen_bias = lm_head_bias.detach() if lm_head_bias is not None else None for step in range(max_horizon): pred_states = pred_states[:, :-1, :] next_token_embeds = next_token_embeds[:, 1:, :] target_states = target_states[:, 1:, :] pred_states = dynamics_model(pred_states, next_token_embeds) # NextLat Eq. 3: SmoothL1(sg[h_{t+i}], hhat_{t+i}) over valid states. mse_mask = valid_tokens[:, step + 1 :] mse_total = mse_total + _masked_smooth_l1( pred_states, target_states, mse_mask, ) if kl_weight > 0.0 and pred_states.shape[1] > 1: # NextLat Eq. 4: KL(p_sg(.|sg[h]) || p_sg(.|hhat)); the LM head is # detached so this auxiliary term trains p_psi/representations, not # the output projection. The standard CCE NTP path updates lm_head. teacher_logits = F.linear( target_states[:, :-1, :].detach(), frozen_weight, frozen_bias, ) student_logits = F.linear( pred_states[:, :-1, :], frozen_weight, frozen_bias, ) kl_mask = valid_tokens[:, step + 2 :] kl_total = kl_total + _categorical_kl_loss( teacher_logits, student_logits, kl_mask, ) inv_horizon = 1.0 / float(max_horizon) mse_loss = mse_total * inv_horizon kl_loss = kl_total * inv_horizon total = float(mse_weight) * mse_loss + float(kl_weight) * kl_loss return total, mse_loss, kl_loss class NeoLLMForCausalLM(NeoLLMPreTrainedModel, GenerationMixin): """ Causal LM with NeoLLM backbone. """ _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"} def __init__(self, config: NeoLLMConfig): super().__init__(config) self.model = NeoLLMModel(config) self.vocab_size = config.vocab_size self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) self.nitp_projector = NITPProjector(config) if config.use_nitp else None self.nitp_temporal_transition = ( NITPTemporalTransition(config) if bool(getattr(config, "use_nitp_temporal", False)) else None ) self.nextlat_dynamics = ( NextLatDynamicsModel(config) if config.use_nextlat else None ) self.ntp_loss_backend = str(getattr(config, "ntp_loss_backend", "cce")).lower() self.liger_ntp_loss = None if self.ntp_loss_backend == "liger": if LigerFusedLinearCrossEntropyLoss is None: raise ImportError( "NeoLLM was configured with `use_liger_kernel=True` / " "`ntp_loss_backend='liger'`, but `liger-kernel` is not " "installed. Install it with `pip install liger-kernel` " "or switch `ntp_loss_backend='cce'`." ) self.liger_ntp_loss = LigerFusedLinearCrossEntropyLoss( ignore_index=-100, reduction="mean", accum_dtype=_resolve_liger_accum_dtype( getattr(config, "liger_loss_accum_dtype", "float32") ), ) elif self.ntp_loss_backend != "cce": raise ValueError( "`ntp_loss_backend` must be 'cce' or 'liger'. " f"Got {self.ntp_loss_backend!r}." ) if bool(getattr(config, "use_mile_loss", False)): _require_extended_cce_options("mile_enabled", "mile_gamma") if bool(getattr(config, "use_mu_loss", False)): _require_extended_cce_options("mu_loss_enabled", "mu_loss_lambda") # The optional MEAP backend is checked only when training actually # applies corruption. Clean training, evaluation, and inference remain # usable without cut-cross-entropy even when loading a MEAP checkpoint. self._last_ntp_loss = None self._last_ntp_ce_unweighted = None self._last_mile_reweighting_delta = None self._last_mu_loss = None self._last_meap_mask_fraction = None self._last_meap_masked_tokens = None self._last_meap_seed = None self._last_tweo_loss = None self._last_nitp_loss = None self._last_nitp_temporal_loss = None self._last_nitp_temporal_state_loss = None self._last_nitp_temporal_step_loss = None self._last_nitp_temporal_raw_step_loss = None self._last_nitp_temporal_mean_error_loss = None self._last_nitp_temporal_std_match_loss = None self._last_nitp_temporal_readout_loss = None self._last_nitp_temporal_identification_loss = None self._last_nitp_temporal_top1_accuracy = None self._last_nitp_temporal_mean_absolute_offset = None self._last_nitp_temporal_identification_margin = None self._last_nitp_temporal_valid_window_fraction = None self._last_nitp_temporal_component_metrics = None self._last_nitp_temporal_step_metrics = None self._last_nitp_temporal_offset_metrics = None self._last_nitp_temporal_scalar_metrics = None self._last_nextlat_loss = None self._last_nextlat_mse_loss = None self._last_nextlat_kl_loss = None self._last_jtokm_aux_loss = None self._last_stack_entropy_loss = None self._last_stack_action_entropy = None self._last_stack_action_entropy_normalized = None self._last_stack_push_prob = None self._last_stack_pop_prob = None self._last_stack_noop_prob = None self._last_stack_action_max_prob = None self._last_stack_expected_depth = None self._last_stack_mask_fill_fraction = None self._last_stack_top_occupancy = None self._last_stack_gate_entropy = None self._last_stack_gate_entropy_normalized = None self._last_stack_gate_max_prob = None self._last_stack_gate_effective_slots = None self._last_total_loss = None if config.use_token_generator: self._tied_weights_keys = {} self.post_init() def get_input_embeddings(self): return self.model.get_input_embeddings() def get_dynamics_metrics(self): """Detached device scalars from the last gradient-enabled training forward. Read outside compiled forward at logging cadence. Sampling describes the last microbatch/recomputation, not an average over optimizer steps. This method does not synchronize, transfer activations, or change loss. """ if not bool(getattr(self.config, "dynamics_metrics_enabled", False)): return {} metrics = {} generator = getattr(self.model, "token_generator", None) packed = getattr(generator, "_last_dynamics_metrics", None) if packed is not None: metrics.update({f"dynamics/leviathan/{name}": value.detach() for name, value in zip(LEV_NAMES, packed.unbind())}) for layer in self.model.layers: module = getattr(layer, "jtok", None) packed = getattr(module, "_last_dynamics_metrics", None) if packed is None: continue kind = "jtokm" if module.use_mixture else "jtok" prefix = f"dynamics/{kind}/layer_{module.layer_idx}/" metrics.update({prefix + name: value.detach() for name, value in unpack_jtok_dynamics(packed, module.num_experts, mixture=module.use_mixture).items()}) return metrics def get_dynamics_parameter_groups(self): """Only investigated generator/bridge/surface/router parameters, no backbone.""" groups = {} generator = getattr(self.model, "token_generator", None) if generator is not None: for name, parameter in generator.named_parameters(): if parameter.requires_grad: groups[f"leviathan/{name}"] = (parameter,) for layer in self.model.layers: module = getattr(layer, "jtok", None) if module is not None: kind = "jtokm" if module.use_mixture else "jtok" for name, parameter in module.named_parameters(): if parameter.requires_grad: groups[f"{kind}/layer_{module.layer_idx}/{name}"] = (parameter,) return groups def set_input_embeddings(self, value): self.model.set_input_embeddings(value) def get_output_embeddings(self): return self.lm_head def set_output_embeddings(self, new_embeddings): self.lm_head = new_embeddings def prepare_inputs_for_generation( self, input_ids: torch.LongTensor, past_key_values=None, attention_mask: Optional[torch.Tensor] = None, inputs_embeds: Optional[torch.FloatTensor] = None, **kwargs, ) -> dict: """ NeoLLM does not implement KV caching (always returns past_key_values=None). Transformers' default GenerationMixin.prepare_inputs_for_generation assumes KV cache is active and slices input_ids to only the newest token on every step past the prefill. Without a real cache that retains previous K/V states, the attention module only sees 1 key/value pair while the causal mask still spans the full context length — causing a shape mismatch in SDPA. This override always forwards the COMPLETE input_ids sequence so the model can recompute attention over the full context from scratch at every step. Generation is therefore slower (no caching benefit) but numerically correct. """ model_inputs: dict = {"input_ids": input_ids, "attention_mask": attention_mask} if inputs_embeds is not None and past_key_values is None: model_inputs["inputs_embeds"] = inputs_embeds return model_inputs def forward( self, input_ids: Optional[torch.LongTensor] = None, attention_mask: Optional[torch.Tensor] = None, position_ids: Optional[torch.LongTensor] = None, inputs_embeds: Optional[torch.FloatTensor] = None, labels: Optional[torch.LongTensor] = None, meap_seed: Optional[Union[int, torch.Tensor]] = None, logits_to_keep: Union[int, torch.Tensor] = 0, output_hidden_states: Optional[bool] = None, return_dict: Optional[bool] = None, **kwargs: Unpack[TransformersKwargs], ) -> CausalLMOutputWithPast: self._last_meap_mask_fraction = None self._last_meap_masked_tokens = None self._last_meap_seed = None meap_enabled = ( bool(getattr(self.config, "use_meap", False)) and self.training and labels is not None ) meap_selected_mask = None if meap_enabled: if input_ids is None or inputs_embeds is not None: raise ValueError( "MEAP requires `input_ids` and does not support precomputed " "`inputs_embeds`." ) if meap_mask_inputs is None: raise ImportError( "Active MEAP training requires the extended cut-cross-entropy " "package. Evaluation and inference do not require it." ) if "return_metrics" not in inspect.signature(meap_mask_inputs).parameters: raise ImportError( "The installed cut-cross-entropy package lacks MEAP metrics. " "Install the extended backend before MEAP training." ) mask_token_id = getattr(self.config, "meap_mask_token_id", None) if mask_token_id is None: raise ValueError("`meap_mask_token_id` is required when MEAP is enabled.") if meap_seed is None: effective_meap_seed = input_ids.new_tensor( int(getattr(self.config, "meap_seed", 0)), dtype=torch.long ) elif isinstance(meap_seed, torch.Tensor): effective_meap_seed = meap_seed.to( device=input_ids.device, dtype=torch.long ) else: effective_meap_seed = input_ids.new_tensor( int(meap_seed), dtype=torch.long ) eligible_mask = ( attention_mask.to(device=input_ids.device, dtype=torch.bool) if attention_mask is not None else None ) exclude_last = bool(getattr(self.config, "meap_exclude_last", True)) mask_ratio = float(getattr(self.config, "meap_mask_ratio", 0.15)) return_meap_mask = bool(getattr(self.config, "use_mile_loss", False)) meap_result = meap_mask_inputs( input_ids, int(mask_token_id), enabled=True, mask_ratio=mask_ratio, eligible_mask=eligible_mask, seed=effective_meap_seed, exclude_last=exclude_last, return_mask=return_meap_mask, return_metrics=True, implementation=str( getattr(self.config, "meap_implementation", "triton") ), ) if return_meap_mask: input_ids, meap_selected_mask, meap_metrics = meap_result else: input_ids, meap_metrics = meap_result eligible_count = meap_metrics[0] masked_count = meap_metrics[1] self._last_meap_mask_fraction = ( masked_count.float() / eligible_count.clamp_min(1).float() ).detach() self._last_meap_masked_tokens = masked_count.detach() self._last_meap_seed = effective_meap_seed.detach() tweo_enabled = ( bool(getattr(self.config, "use_tweo", False)) and labels is not None and float(getattr(self.config, "tweo_loss_weight", 0.0)) > 0.0 ) nitp_temporal_enabled = ( bool(getattr(self.config, "use_nitp_temporal", False)) and self.nitp_temporal_transition is not None and self.nitp_projector is not None and labels is not None ) nitp_enabled = ( bool(getattr(self.config, "use_nitp", False)) and self.nitp_projector is not None and labels is not None and ( float(getattr(self.config, "nitp_loss_weight", 0.0)) > 0.0 or nitp_temporal_enabled ) ) nextlat_enabled = ( bool(getattr(self.config, "use_nextlat", False)) and self.nextlat_dynamics is not None and labels is not None and ( float(getattr(self.config, "nextlat_mse_weight", 0.0)) > 0.0 or float(getattr(self.config, "nextlat_kl_weight", 0.0)) > 0.0 ) ) model_output_hidden_states = ( True if (nitp_enabled or nitp_temporal_enabled or nextlat_enabled) else output_hidden_states ) model_out = self.model( input_ids=input_ids, attention_mask=attention_mask, position_ids=position_ids, inputs_embeds=inputs_embeds, output_hidden_states=model_output_hidden_states, output_tweo_activations=tweo_enabled, meap_active=meap_enabled, return_dict=return_dict, **kwargs, ) # Unpack: model returns BaseModelOutputWithPast, or a tuple when # return_dict=False (with compact JTok-M stats and/or tweo activations # appended when enabled). tweo_activations = None hidden_states_tuple = None jtokm_aux_stats = None stack_metrics = None if isinstance(model_out, tuple): outputs = model_out[0] if tweo_enabled and len(model_out) >= 2: tweo_activations = model_out[-1] if isinstance(outputs, torch.Tensor): hidden_states = outputs if model_output_hidden_states: tuple_candidates = [ candidate for candidate in model_out[1:] if isinstance(candidate, tuple) and (not candidate or not isinstance(candidate[0], tuple)) ] if tuple_candidates: hidden_states_tuple = tuple_candidates[0] for candidate in model_out[1:]: if ( bool(getattr(self.config, "use_stack_memory", False)) and isinstance(candidate, torch.Tensor) and candidate.ndim == 2 and candidate.shape[-1] == 11 ): stack_metrics = candidate if ( isinstance(candidate, tuple) and candidate and isinstance(candidate[0], tuple) and len(candidate[0]) == 3 ): jtokm_aux_stats = candidate elif isinstance(outputs, tuple): hidden_states = outputs[0] if len(outputs) > 2: hidden_states_tuple = outputs[2] else: hidden_states = outputs.last_hidden_state hidden_states_tuple = outputs.hidden_states jtokm_aux_stats = getattr(outputs, "jtokm_aux_stats", None) stack_metrics = getattr(outputs, "stack_metrics", None) else: outputs = model_out hidden_states = outputs.last_hidden_state hidden_states_tuple = outputs.hidden_states jtokm_aux_stats = getattr(outputs, "jtokm_aux_stats", None) stack_metrics = getattr(outputs, "stack_metrics", None) loss = None ntp_loss = None tweo_loss = None nitp_loss = None nitp_temporal_loss = None nitp_temporal_state_loss = None nitp_temporal_step_loss = None nitp_temporal_readout_loss = None nitp_temporal_identification_loss = None nitp_temporal_top1_accuracy = None nitp_temporal_mean_absolute_offset = None nitp_temporal_identification_margin = None nitp_temporal_valid_window_fraction = None nitp_temporal_component_metrics = None nitp_temporal_step_metrics = None nitp_temporal_offset_metrics = None nitp_temporal_scalar_metrics = None nitp_target_hidden_states = None nextlat_loss = None nextlat_mse_loss = None nextlat_kl_loss = None jtokm_aux_loss = None stack_entropy_loss = None self._last_ntp_loss = None self._last_ntp_ce_unweighted = None self._last_mile_reweighting_delta = None self._last_mu_loss = None self._last_tweo_loss = None self._last_nitp_loss = None self._last_nitp_temporal_loss = None self._last_nitp_temporal_state_loss = None self._last_nitp_temporal_step_loss = None self._last_nitp_temporal_raw_step_loss = None self._last_nitp_temporal_mean_error_loss = None self._last_nitp_temporal_std_match_loss = None self._last_nitp_temporal_readout_loss = None self._last_nitp_temporal_identification_loss = None self._last_nitp_temporal_top1_accuracy = None self._last_nitp_temporal_mean_absolute_offset = None self._last_nitp_temporal_identification_margin = None self._last_nitp_temporal_valid_window_fraction = None self._last_nitp_temporal_component_metrics = None self._last_nitp_temporal_step_metrics = None self._last_nitp_temporal_offset_metrics = None self._last_nitp_temporal_scalar_metrics = None self._last_nextlat_loss = None self._last_nextlat_mse_loss = None self._last_nextlat_kl_loss = None self._last_jtokm_aux_loss = None self._last_stack_entropy_loss = None self._last_stack_action_entropy = None self._last_stack_action_entropy_normalized = None self._last_stack_push_prob = None self._last_stack_pop_prob = None self._last_stack_noop_prob = None self._last_stack_action_max_prob = None self._last_stack_expected_depth = None self._last_stack_mask_fill_fraction = None self._last_stack_top_occupancy = None self._last_stack_gate_entropy = None self._last_stack_gate_entropy_normalized = None self._last_stack_gate_max_prob = None self._last_stack_gate_effective_slots = None self._last_total_loss = None if labels is not None: if self.ntp_loss_backend == "liger": ntp_loss = compute_liger_loss( hidden_states, labels, self.lm_head.weight, getattr(self.lm_head, "bias", None), self.config.pad_token_id, self.liger_ntp_loss, ) else: return_loss_metrics = bool( getattr(self.config, "use_mile_loss", False) or getattr(self.config, "use_mu_loss", False) or meap_enabled ) ntp_result = compute_cce_loss( hidden_states, labels, self.lm_head.weight, getattr(self.lm_head, "bias", None), self.config.pad_token_id, getattr(self.config, "cce_loss_impl", "cce_kahan_full_c"), use_mile_loss=bool( getattr(self.config, "use_mile_loss", False) ), mile_loss_gamma=float( getattr(self.config, "mile_loss_gamma", 1.0) ), mile_group_mask=( meap_selected_mask if meap_enabled and bool(getattr(self.config, "use_mile_loss", False)) else None ), use_mu_loss=bool(getattr(self.config, "use_mu_loss", False)), mu_loss_lambda=float( getattr(self.config, "mu_loss_lambda", 1e-4) ), return_loss_metrics=return_loss_metrics, ) if return_loss_metrics: ntp_loss, ntp_metrics = ntp_result self._last_ntp_ce_unweighted = ntp_metrics[ "ntp_ce_unweighted" ].detach() self._last_mile_reweighting_delta = ntp_metrics[ "mile_reweighting_delta" ].detach() self._last_mu_loss = ntp_metrics["mu_loss"].detach() else: ntp_loss = ntp_result loss = ntp_loss # StackTrans entropy regularizer. The backbone returns one compact # metric row per decoder layer; column 0 is the differentiable mean # entropy H(push,pop,noop). Averaging across layers keeps lambda # invariant to model depth, batch size, sequence length, and head count. if bool(getattr(self.config, "use_stack_memory", False)): if stack_metrics is None or stack_metrics.numel() == 0: raise RuntimeError( "StackMemory is active but no stack metrics were returned." ) stack_action_entropy = stack_metrics[:, 0].mean() stack_entropy_weight = torch.as_tensor( float(getattr(self.config, "stack_entropy_loss_weight", 1e-3)), device=stack_action_entropy.device, dtype=torch.float32, ) stack_entropy_loss = stack_entropy_weight * stack_action_entropy loss = loss + stack_entropy_loss if tweo_enabled: if not tweo_activations: raise ValueError( "`use_tweo=True` requires Transformer block activations. " "NeoLLMForCausalLM.forward enables them automatically " "while labels are present." ) tweo_loss = compute_tweo_loss( block_activations=tweo_activations, tau=self.config.tweo_tau, power=self.config.tweo_power, eps=self.config.tweo_eps, attention_mask=attention_mask, ) # Keep the loss coefficient on the same CUDA graph as the loss. # A Python float is lifted as a CPU scalar by AOTAutograd and can # invalidate CUDAGraph capture during the compiled backward pass. tweo_loss_weight = torch.as_tensor( self.config.tweo_loss_weight, device=tweo_loss.device, dtype=torch.float32, ) loss = loss + tweo_loss_weight * tweo_loss if nitp_enabled: target_layer = int(getattr(self.config, "nitp_target_layer", 1)) if ( hidden_states_tuple is None or len(hidden_states_tuple) <= target_layer ): raise ValueError( "`use_nitp=True` requires backbone hidden states including " f"target layer {target_layer}. " "Call the model with output_hidden_states=True or leave " "NeoLLMForCausalLM.forward to enable it automatically." ) nitp_loss = compute_nitp_loss( final_hidden_states=hidden_states, target_hidden_states=hidden_states_tuple[target_layer], labels=labels, attention_mask=attention_mask, projector=self.nitp_projector, pad_token_id=self.config.pad_token_id, ) nitp_target_hidden_states = hidden_states_tuple[target_layer] if float(self.config.nitp_loss_weight) > 0.0: nitp_loss_weight = torch.as_tensor( self.config.nitp_loss_weight, device=nitp_loss.device, dtype=torch.float32, ) loss = loss + nitp_loss_weight * nitp_loss if nitp_temporal_enabled: if hidden_states_tuple is None or len(hidden_states_tuple) < 2: raise ValueError( "`use_nitp_temporal=True` requires backbone hidden " "states. NeoLLMForCausalLM.forward enables them " "automatically while labels are present." ) if nitp_target_hidden_states is None: target_layer = int(getattr(self.config, "nitp_target_layer", 1)) if len(hidden_states_tuple) <= target_layer: raise ValueError( "`use_nitp_temporal=True` requires the NITP shallow " f"target layer {target_layer}." ) nitp_target_hidden_states = hidden_states_tuple[target_layer] temporal_apply_loss = bool( getattr(self.config, "nitp_temporal_apply_loss", True) ) eos_token_id = ( self.config.eos_token_id[0] if isinstance(self.config.eos_token_id, (list, tuple)) and self.config.eos_token_id else self.config.eos_token_id ) if temporal_apply_loss: temporal_results = compute_nitp_temporal_objective( hidden_states=hidden_states, shallow_target_states=nitp_target_hidden_states, labels=labels, attention_mask=attention_mask, transition_model=self.nitp_temporal_transition, nitp_projector=self.nitp_projector, lm_head_weight=self.lm_head.weight, lm_head_bias=getattr(self.lm_head, "bias", None), horizon=self.config.nitp_temporal_horizon, dynamics_weight=self.config.nitp_temporal_dynamics_weight, identification_weight=( self.config.nitp_temporal_identification_weight ), similarity_temperature=( self.config.nitp_temporal_similarity_temperature ), target_temperature=( self.config.nitp_temporal_target_temperature ), pad_token_id=self.config.pad_token_id, eos_token_id=eos_token_id, ) else: # Monitoring-only mode: the branch is instantiated (and # therefore contributes parameters) but its graph is not # retained and its objective is not added to training. with torch.no_grad(): temporal_results = compute_nitp_temporal_objective( hidden_states=hidden_states, shallow_target_states=nitp_target_hidden_states, labels=labels, attention_mask=attention_mask, transition_model=self.nitp_temporal_transition, nitp_projector=self.nitp_projector, lm_head_weight=self.lm_head.weight, lm_head_bias=getattr(self.lm_head, "bias", None), horizon=self.config.nitp_temporal_horizon, dynamics_weight=( self.config.nitp_temporal_dynamics_weight ), identification_weight=( self.config.nitp_temporal_identification_weight ), similarity_temperature=( self.config.nitp_temporal_similarity_temperature ), target_temperature=( self.config.nitp_temporal_target_temperature ), pad_token_id=self.config.pad_token_id, eos_token_id=eos_token_id, ) ( nitp_temporal_loss, nitp_temporal_state_loss, nitp_temporal_step_loss, nitp_temporal_readout_loss, nitp_temporal_identification_loss, nitp_temporal_top1_accuracy, nitp_temporal_mean_absolute_offset, nitp_temporal_identification_margin, nitp_temporal_valid_window_fraction, nitp_temporal_component_metrics, nitp_temporal_step_metrics, nitp_temporal_offset_metrics, nitp_temporal_scalar_metrics, ) = temporal_results if temporal_apply_loss: loss = loss + nitp_temporal_loss if nextlat_enabled: if hidden_states_tuple is None or len(hidden_states_tuple) < 2: raise ValueError( "`use_nextlat=True` requires backbone hidden states. " "NeoLLMForCausalLM.forward enables them automatically " "while labels are present." ) # NextLat keeps its original teacher-forced token-embedding # path. This assignment is local to NextLat; causal Temporal # NITP no longer reads ``hidden_states_tuple[0]``. token_embeds = hidden_states_tuple[0] nextlat_loss, nextlat_mse_loss, nextlat_kl_loss = compute_nextlat_loss( hidden_states=hidden_states, token_embeds=token_embeds, labels=labels, attention_mask=attention_mask, dynamics_model=self.nextlat_dynamics, lm_head_weight=self.lm_head.weight, lm_head_bias=getattr(self.lm_head, "bias", None), horizon=self.config.nextlat_horizon, mse_weight=self.config.nextlat_mse_weight, kl_weight=self.config.nextlat_kl_weight, pad_token_id=self.config.pad_token_id, ) loss = loss + nextlat_loss if bool(getattr(self.config, "use_jtokm", False)) and jtokm_aux_stats: jtokm_aux_loss = compute_jtokm_aux_loss( jtokm_aux_stats, n_e=int(getattr(self.config, "jtok_num_experts", 5)), top_k=int(getattr(self.config, "jtok_top_k", 2)), weight=float(getattr(self.config, "jtok_aux_loss_weight", 1e-4)), ) loss = loss + jtokm_aux_loss self._last_ntp_loss = ntp_loss.detach() if ntp_loss is not None else None self._last_tweo_loss = tweo_loss.detach() if tweo_loss is not None else None self._last_nitp_loss = nitp_loss.detach() if nitp_loss is not None else None self._last_nitp_temporal_loss = ( nitp_temporal_loss.detach() if nitp_temporal_loss is not None else None ) self._last_nitp_temporal_state_loss = ( nitp_temporal_state_loss.detach() if nitp_temporal_state_loss is not None else None ) self._last_nitp_temporal_step_loss = ( nitp_temporal_step_loss.detach() if nitp_temporal_step_loss is not None else None ) self._last_nitp_temporal_readout_loss = ( nitp_temporal_readout_loss.detach() if nitp_temporal_readout_loss is not None else None ) self._last_nitp_temporal_identification_loss = ( nitp_temporal_identification_loss.detach() if nitp_temporal_identification_loss is not None else None ) self._last_nitp_temporal_top1_accuracy = ( nitp_temporal_top1_accuracy.detach() if nitp_temporal_top1_accuracy is not None else None ) self._last_nitp_temporal_mean_absolute_offset = ( nitp_temporal_mean_absolute_offset.detach() if nitp_temporal_mean_absolute_offset is not None else None ) self._last_nitp_temporal_identification_margin = ( nitp_temporal_identification_margin.detach() if nitp_temporal_identification_margin is not None else None ) self._last_nitp_temporal_valid_window_fraction = ( nitp_temporal_valid_window_fraction.detach() if nitp_temporal_valid_window_fraction is not None else None ) self._last_nitp_temporal_component_metrics = ( nitp_temporal_component_metrics.detach() if nitp_temporal_component_metrics is not None else None ) self._last_nitp_temporal_step_metrics = ( nitp_temporal_step_metrics.detach() if nitp_temporal_step_metrics is not None else None ) # Aggregate component metrics are derived from the append-only rows # 13..15. The existing `_last_nitp_temporal_step_loss` remains the # composite raw + mean_error + 0.5*std value for compatibility. if ( nitp_temporal_step_metrics is not None and nitp_temporal_step_metrics.shape[0] >= 16 and nitp_temporal_step_metrics.shape[1] > 0 ): detached_step_metrics = nitp_temporal_step_metrics.detach() self._last_nitp_temporal_raw_step_loss = ( detached_step_metrics[13].mean() ) self._last_nitp_temporal_mean_error_loss = ( detached_step_metrics[14].mean() ) self._last_nitp_temporal_std_match_loss = ( detached_step_metrics[15].mean() ) self._last_nitp_temporal_offset_metrics = ( nitp_temporal_offset_metrics.detach() if nitp_temporal_offset_metrics is not None else None ) self._last_nitp_temporal_scalar_metrics = ( nitp_temporal_scalar_metrics.detach() if nitp_temporal_scalar_metrics is not None else None ) self._last_nextlat_loss = ( nextlat_loss.detach() if nextlat_loss is not None else None ) self._last_nextlat_mse_loss = ( nextlat_mse_loss.detach() if nextlat_mse_loss is not None else None ) self._last_nextlat_kl_loss = ( nextlat_kl_loss.detach() if nextlat_kl_loss is not None else None ) self._last_jtokm_aux_loss = ( jtokm_aux_loss.detach() if jtokm_aux_loss is not None else None ) if bool(getattr(self.config, "use_stack_memory", False)): detached_stack = stack_metrics.detach().float() action_entropy = detached_stack[:, 0].mean() gate_entropy = detached_stack[:, 8].mean() self._last_stack_entropy_loss = ( stack_entropy_loss.detach() if stack_entropy_loss is not None else action_entropy.new_zeros(()) ) self._last_stack_action_entropy = action_entropy self._last_stack_action_entropy_normalized = action_entropy / math.log(3.0) self._last_stack_push_prob = detached_stack[:, 1].mean() self._last_stack_pop_prob = detached_stack[:, 2].mean() self._last_stack_noop_prob = detached_stack[:, 3].mean() self._last_stack_action_max_prob = detached_stack[:, 4].mean() self._last_stack_expected_depth = detached_stack[:, 5].mean() self._last_stack_mask_fill_fraction = detached_stack[:, 6].mean() self._last_stack_top_occupancy = detached_stack[:, 7].mean() self._last_stack_gate_entropy = gate_entropy self._last_stack_gate_entropy_normalized = gate_entropy / math.log( float(self.config.stack_slots) ) self._last_stack_gate_max_prob = detached_stack[:, 9].mean() self._last_stack_gate_effective_slots = detached_stack[:, 10].mean() self._last_total_loss = loss.detach() if loss is not None else None logits = None else: slice_indices = ( slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep ) logits = self.lm_head(hidden_states[:, slice_indices, :]) return CausalLMOutputWithPast( loss=loss, logits=logits, past_key_values=None, hidden_states=hidden_states_tuple, attentions=outputs.attentions if hasattr(outputs, "attentions") else None, ) # ==================== AUTOMODEL REGISTRATION ==================== __all__ = [ "NeoLLMForCausalLM", "NeoLLMModel", "NeoLLMPreTrainedModel", "NeoLLMConfig", "LeviathanGenerator", "LeviathanJTok", "NeoLLMBaseModelOutputWithPast", "compute_jtokm_aux_loss", "SpellingBeeEmbedding", "StackMemory", "FANLayer", "ScalarMultiplier", "VectorMultiplier", "LinearWithMultipliers", "MEAHeadRMSNorm", "HadamardOProj", "REPOModule", "RepoGrapePositioning", "RepoGoatPrior", ] AutoConfig.register("neollm", NeoLLMConfig) AutoModel.register(NeoLLMConfig, NeoLLMModel) AutoModelForCausalLM.register(NeoLLMConfig, NeoLLMForCausalLM)