# Copyright (c) 2026, the ComplexKDA authors. # Copyright (c) 2023-2026, Songlin Yang, Yu Zhang, Zhiyuan Li (the parts derived # from flash-linear-attention, MIT licensed). # # SPDX-License-Identifier: MIT """ComplexKDA -- Kimi Delta Attention with a signed (Z_2-phased) decay gate. STANDALONE. This file needs only `torch` and `transformers`. It carries its own implementations of everything the model is made of -- the short convolution, the RMS norms, the SwiGLU MLP, the attention layers of the hybrid arms, and the gated-delta recurrence itself -- so a checkpoint loads and runs with nothing else installed. IT GOES FASTER WITH THE FORK. When `fla` from https://github.com/OpenEuroLLM/ComplexKDA is importable, the Triton kernels and fused modules it ships are used instead, and the model is then running exactly the code the checkpoints were trained through. Detection is by capability, not by name: upstream flash-linear-attention also provides `chunk_kda`, but without the `sign` argument the signed gate needs, so the signature is inspected rather than trusted. Set the environment variable `COMPLEX_KDA_BACKEND=torch` to force the pure-torch path (useful for debugging a numerical difference), or `=kernel` to make a missing fork an error rather than a silent fallback. WHAT THE SIGNED GATE IS. A gated-delta layer carries a per-channel decay `alpha`; KDA, like every gated linear attention before it, confines it to `(0, 1]`. ComplexKDA lets it take either sign, `alpha in [-1, 1]`, which is the one-dimensional real case of a complex eigenvalue -- a channel can now oscillate rather than only forget. The magnitude is carried in log space exactly as before, and the `+-1` part is carried separately as a running product (the "gauge") pushed onto q and k, so the recurrence the kernels run is still the unsigned one. `running_sign` below is that product, and `ungauge_state` takes it back off the state at a chunk boundary so a cached state is the real one. """ from __future__ import annotations import math import os import warnings from typing import Any import torch import torch.nn as nn import torch.nn.functional as F from transformers.generation import GenerationMixin from transformers.modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast from transformers.modeling_utils import PreTrainedModel from transformers.utils import logging try: # packaged next to the weights (the Hub layout) from .configuration_complex_kda import ComplexKDAConfig, get_hybrid_attention_spec except ImportError: # imported as a loose file from configuration_complex_kda import ComplexKDAConfig, get_hybrid_attention_spec logger = logging.get_logger(__name__) __all__ = [ "ComplexKDACache", "ComplexKDAForCausalLM", "ComplexKDAModel", "ComplexKDAPreTrainedModel", "ComplexKimiDeltaAttention", ] # =========================================================================== # optional fast path # =========================================================================== def _detect_fla(): """(chunk_kda, fused_recurrent_kda) from the fork, or (None, None). The test is that the op ACCEPTS `sign`. Upstream fla exports a `chunk_kda` of the same name that computes the unsigned recurrence; calling it with a signed checkpoint's weights would return a plausible tensor rather than an error, so the presence of the module is not enough. """ try: import inspect from fla.ops.kda import chunk_kda, fused_recurrent_kda except Exception: return None, None try: if "sign" not in inspect.signature(chunk_kda).parameters: return None, None if "sign" not in inspect.signature(fused_recurrent_kda).parameters: return None, None except (TypeError, ValueError): return None, None return chunk_kda, fused_recurrent_kda _CHUNK_KDA, _FUSED_RECURRENT_KDA = _detect_fla() def _triton_launchable() -> bool: """Stricter than importing triton: fla imports fine on a CPU-only box and only fails when a kernel is launched.""" try: import triton # noqa: F401 return torch.cuda.is_available() except Exception: return False HAS_KERNEL = _CHUNK_KDA is not None and _triton_launchable() _REQUESTED = os.environ.get("COMPLEX_KDA_BACKEND", "auto").lower() if _REQUESTED not in ("auto", "torch", "kernel"): raise ValueError(f"COMPLEX_KDA_BACKEND must be 'auto', 'torch' or 'kernel'; got {_REQUESTED!r}") if _REQUESTED == "kernel" and not HAS_KERNEL: raise ImportError( "COMPLEX_KDA_BACKEND=kernel, but the ComplexKDA fla fork's signed kernels are not " "available (need a CUDA device, triton, and `pip install " "git+https://github.com/OpenEuroLLM/ComplexKDA`).") USE_KERNEL = HAS_KERNEL and _REQUESTED != "torch" # Chunk length of the portable recurrence. It trades memory for sequential # steps: the intra-chunk term materialises a [chunk, chunk, head_dim] block per # head, so 64 is a few tens of MB at these geometries and 256 is a few hundred. # It does not change what is computed -- only the order the same sums are taken # in, which at fp32 moves a logit by ~1e-6 per layer. CHUNK_SIZE = int(os.environ.get("COMPLEX_KDA_CHUNK_SIZE", "64")) if CHUNK_SIZE <= 0: raise ValueError(f"COMPLEX_KDA_CHUNK_SIZE must be positive; got {CHUNK_SIZE}") if not USE_KERNEL: logger.warning_once( "ComplexKDA is running its portable torch implementation. For the Triton kernels the " "models were trained with, install the fork: " "`pip install git+https://github.com/OpenEuroLLM/ComplexKDA` (and set " "COMPLEX_KDA_BACKEND=torch to keep this path).") def _fla_modules(): """fla's fused ShortConvolution / RMSNorm / gated RMSNorm, or (None,)*3.""" if not USE_KERNEL: return None, None, None try: from fla.modules import FusedRMSNormGated, RMSNorm, ShortConvolution return ShortConvolution, RMSNorm, FusedRMSNormGated except Exception: return None, None, None _FLA_SHORTCONV, _FLA_RMSNORM, _FLA_RMSNORM_GATED = _fla_modules() # =========================================================================== # the gate # # name alpha range activation # "softplus" (0, 1] -exp(A_log) * softplus(u) # "sigmoid" (0, 1] lower_bound * sigmoid(A * u) # "signed_sigmoid2" [-1, 1] 2*sigmoid(u) - 1, evaluated as tanh(u/2) # "signed_tanh" [-1, 1] tanh(u) # # The published baselines use "sigmoid"; the ComplexKDA arms use # "signed_sigmoid2". # =========================================================================== GATES = ("softplus", "sigmoid", "signed_sigmoid2", "signed_tanh") def is_signed(gate: str) -> bool: if gate not in GATES: raise ValueError(f"gate must be one of {GATES}, got {gate!r}") return gate.startswith("signed_") def safe_gate_ok(gate: str) -> bool: """Whether log|alpha| is bounded below by `lower_bound`. False only for "softplus", which is unbounded.""" return gate != "softplus" def signed_gate(z, A_log=None, dt_bias=None, lower_bound: float = -5.0, activation: str = "sigmoid2"): """One pre-activation -> (sign, log|alpha|), for alpha in [-1, 1]. |alpha| = eps + (1 - eps) * |a|, eps = exp(lower_bound), a = tanh(u/2) or tanh(u). "sigmoid2" is `2*sigmoid(u) - 1` written as `tanh(u/2)`: the literal spelling cancels catastrophically near u = 0, where the SIGN is decided, so it would be settled by rounding rather than by u. The sign shares z with the magnitude and is locally constant, so detaching it is exact. """ eps = math.exp(lower_bound) u = z.float() if dt_bias is not None: u = u + dt_bias.view(*([1] * (z.dim() - 2)), *z.shape[-2:]) if A_log is not None: u = A_log.float().exp().view(*([1] * (z.dim() - 2)), -1, 1) * u if activation == "sigmoid2": a = torch.tanh(0.5 * u) elif activation == "tanh": a = torch.tanh(u) else: raise ValueError(f"activation must be 'sigmoid2' or 'tanh', got {activation!r}") s = torch.where(a.detach() < 0, -1, 1).to(torch.int8) return s, (eps + (1.0 - eps) * a.abs()).log() def compute_gate(gate, z, A_log=None, dt_bias=None, lower_bound=-5.0): """name -> (sign int8 or None, log|alpha| fp32). A None sign is what tells the caller there is no gauge to apply.""" if gate not in GATES: raise ValueError(f"gate must be one of {GATES}, got {gate!r}") if gate.startswith("signed_"): return signed_gate(z, A_log, dt_bias, lower_bound, activation=gate[len("signed_"):]) u = z.float() if dt_bias is not None: u = u + dt_bias.view(*([1] * (z.dim() - 2)), *z.shape[-2:]) A = A_log.float().exp().view(*([1] * (z.dim() - 2)), -1, 1) if A_log is not None else 1.0 if gate == "softplus": return None, -A * F.softplus(u) return None, lower_bound * torch.sigmoid(A * u) def signed_gate_init(dt, lower_bound: float = -5.0, activation: str = "sigmoid2"): """dt_bias giving alpha = +exp(-dt) at step 0.""" eps = math.exp(lower_bound) target = ((torch.exp(-dt) - eps) / (1 - eps)).clamp(1e-7, 1 - 1e-7) inv = torch.atanh(target) return 2.0 * inv if activation == "sigmoid2" else inv def gate_init(gate, dt, lower_bound=-5.0): """dt_bias init inverting each gate's own forward, so all four gates start at the same alpha = exp(-dt).""" if gate.startswith("signed_"): return signed_gate_init(dt, lower_bound, gate[len("signed_"):]) if gate == "sigmoid": p = (dt / abs(lower_bound)).clamp(1e-7, 1 - 1e-7) return torch.log(p) - torch.log1p(-p) return dt + torch.log(-torch.expm1(-dt)) def init_dt_bias(gate, gate_dim=None, lower_bound=-5.0, gate_init_style="shipped", dt=None): if dt is None: dt = torch.exp( torch.rand(gate_dim, dtype=torch.float32) * (math.log(0.1) - math.log(0.001)) + math.log(0.001) ).clamp(min=1e-4) init = gate_init(gate, dt, lower_bound) if is_signed(gate) and gate_init_style == "spread": # Same |alpha| as "shipped" with the sign flipped on half the channels: # the activation is odd, so sign and magnitude do not trade off. init = init * torch.where(torch.rand_like(init) < 0.5, -1.0, 1.0) return init # =========================================================================== # the gauge: carry the +-1 part of alpha as a running sign on q/k # =========================================================================== def running_sign(s: torch.Tensor, cu_seqlens: torch.Tensor | None = None) -> torch.Tensor: """P_t = prod_{u<=t} s_u along dim 1, as int8. An integer parity prefix sum: exact at any length, and it carries no autograd graph, because the sign has no gradient. Resets at sequence starts when `cu_seqlens` is given. """ bits = (s < 0).to(torch.int32) par = bits.cumsum(dim=1) if cu_seqlens is not None: starts = cu_seqlens[:-1] idx = torch.repeat_interleave(starts, cu_seqlens[1:] - starts) par = par - (par[:, idx] - bits[:, idx]) return torch.where(par & 1 == 1, -1, 1).to(torch.int8) class _ApplySign(torch.autograd.Function): """x * P, keeping P as int8 rather than letting `mul` upcast it.""" @staticmethod def forward(ctx, x, P): ctx.save_for_backward(P) return x * P.to(x.dtype) @staticmethod def backward(ctx, go): (P,) = ctx.saved_tensors return go * P.to(go.dtype), None def apply_sign(x, P): return _ApplySign.apply(x, P) def ungauge_state(ht, P_last, state_v_first: bool, head_k_dim: int | None = None): """S_T = Diag(P_T) S~_T, on whichever axis holds K. `state_v_first=True` stores [N, HV, V, K] -- K LAST -- so the axis differs between the kernel and the torch path; `head_k_dim` turns a silent wrong-axis bug into an assert. """ if ht is None or P_last is None: return ht if P_last.ndim != ht.ndim - 1: raise AssertionError( f"gauge rank mismatch: state {tuple(ht.shape)} takes a gauge of " f"{ht.ndim - 1} dims, got {tuple(P_last.shape)}.") axis = -1 if state_v_first else -2 if head_k_dim is not None and ht.shape[axis] != head_k_dim: raise AssertionError( f"state layout mismatch: state_v_first={state_v_first} implies K on axis {axis}, " f"but state shape {tuple(ht.shape)} has {ht.shape[axis]} there, not " f"head_k_dim={head_k_dim}.") P = P_last.to(ht.dtype) return ht * (P.unsqueeze(-2) if state_v_first else P.unsqueeze(-1)) # =========================================================================== # the recurrence, in torch # # Both functions take q/k ALREADY l2-normalised and gauged, `g` as log|alpha|, # and `beta` already through its sigmoid -- the same contract as the reference # implementation in the fork, so the two can be compared term by term. # =========================================================================== def recurrent_kda_torch( q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, g: torch.Tensor, beta: torch.Tensor, scale: float | None = None, initial_state: torch.Tensor | None = None, output_final_state: bool = False, ): """The definition, one step at a time. [B,T,H,K] q/k, [B,T,HV,V] v, [B,T,HV,K] g, [B,T,HV] beta; state [B,HV,K,V].""" dtype = v.dtype B, T, H, K = q.shape HV, V = v.shape[2], v.shape[-1] G = HV // H if scale is None: scale = K ** -0.5 q, k, v, g, beta = (x.float() for x in (q, k, v, g, beta)) q = q.repeat_interleave(G, dim=2) * scale k = k.repeat_interleave(G, dim=2) S = q.new_zeros(B, HV, K, V) if initial_state is not None: S = S + initial_state.float() o = torch.zeros_like(v) for i in range(T): q_i, k_i, v_i, g_i, b_i = q[:, i], k[:, i], v[:, i], g[:, i], beta[:, i] S = S * g_i[..., None].exp() # delta rule: replace the memory currently read out by k_i with v_i S = S + torch.einsum("bhk,bhv->bhkv", b_i[..., None] * k_i, v_i - (k_i[..., None] * S).sum(-2)) o[:, i] = torch.einsum("bhk,bhkv->bhv", q_i, S) return o.to(dtype), (S if output_final_state else None) def chunk_kda_torch( q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, g: torch.Tensor, beta: torch.Tensor, scale: float | None = None, initial_state: torch.Tensor | None = None, output_final_state: bool = False, chunk_size: int = 64, ): """The same recurrence in chunks: O(T/C) sequential steps instead of O(T). The WY/UT transform of the chunk's delta updates, then one state carry per chunk. Arithmetically identical to `recurrent_kda_torch` up to floating point; it exists because a 4096-token forward through the step loop is minutes rather than milliseconds. MASK BEFORE EXPONENTIATING. Every exponent used here is a sum of `log|alpha|` over an interval, so it is <= 0 and `exp` is safe -- but only for the pairs the causal mask keeps. The reference implementation exponentiates the full block and masks afterwards, which overflows once `|log alpha| * chunk` passes ~88 in fp32: finite forward, NaN backward. Masking first removes that failure mode entirely, which is why `chunk_size` needs no upper bound here. """ dtype = v.dtype B, T, H, K = q.shape HV, V = v.shape[2], v.shape[-1] G = HV // H if scale is None: scale = K ** -0.5 BT = int(chunk_size) if BT <= 0: raise ValueError(f"chunk_size must be positive, got {chunk_size}") q, k, v, g, beta = (x.float() for x in (q, k, v, g, beta)) q = q.repeat_interleave(G, dim=2) * scale k = k.repeat_interleave(G, dim=2) # Pad the tail to a whole chunk. beta = 0 makes the padded steps write # nothing and g = 0 makes them decay nothing, so the carried state is # exactly the state at T. pad = (-T) % BT if pad: q = F.pad(q, (0, 0, 0, 0, 0, pad)) k = F.pad(k, (0, 0, 0, 0, 0, pad)) v = F.pad(v, (0, 0, 0, 0, 0, pad)) g = F.pad(g, (0, 0, 0, 0, 0, pad)) beta = F.pad(beta, (0, 0, 0, pad)) NT = (T + pad) // BT # [B, T, HV, X] -> [B, HV, NT, BT, X] def _chunks(x): return x.view(B, NT, BT, *x.shape[2:]).permute(0, 3, 1, 2, *range(4, x.dim() + 1)) q, k, v, g = (_chunks(x) for x in (q, k, v, g)) beta = beta.view(B, NT, BT, HV).permute(0, 3, 1, 2) eye = torch.eye(BT, device=q.device, dtype=q.dtype) rows = torch.arange(BT, device=q.device) strictly_lower = rows[:, None] > rows[None, :] # c > i causal = rows[:, None] >= rows[None, :] # c >= j neg_inf = torch.finfo(q.dtype).min S = q.new_zeros(B, HV, K, V) if initial_state is not None: S = S + initial_state.float() o = torch.zeros_like(v) for n in range(NT): q_n, k_n, v_n, g_n, b_n = q[:, :, n], k[:, :, n], v[:, :, n], g[:, :, n], beta[:, :, n] gc = g_n.cumsum(-2) # [B,HV,BT,K], <= 0 # T[c,i] = beta_c * for c > i -- the # strictly-lower part of the chunk's own delta interactions. d = gc.unsqueeze(-2) - gc.unsqueeze(-3) # [B,HV,BT(c),BT(i),K] d = d.masked_fill(~strictly_lower[..., None], neg_inf) A = (k_n.unsqueeze(-2) * d.exp() * k_n.unsqueeze(-3)).sum(-1) del d A = -(A * b_n[..., :, None]) # (I - A)^{-1}, A strictly lower and hence unit-triangular after +I. # The reference walks the Neumann series row by row; a triangular solve # is the same matrix and vectorises. Ainv = torch.linalg.solve_triangular(eye - A, eye.expand_as(A), upper=False, unitriangular=True) Aw = Ainv * b_n[..., None, :] w = Aw @ (gc.exp() * k_n) # [B,HV,BT,K] u = Aw @ v_n # [B,HV,BT,V] dq = gc.unsqueeze(-2) - gc.unsqueeze(-3) dq = dq.masked_fill(~causal[..., None], neg_inf) Aqk = (q_n.unsqueeze(-2) * dq.exp() * k_n.unsqueeze(-3)).sum(-1) del dq v_new = u - w @ S o[:, :, n] = (q_n * gc.exp()) @ S + Aqk @ v_new g_last = gc[:, :, -1] # [B,HV,K] S = S * g_last.unsqueeze(-1).exp() S = S + ((g_last.unsqueeze(-2) - gc).exp() * k_n).transpose(-1, -2) @ v_new o = o.permute(0, 2, 3, 1, 4).reshape(B, NT * BT, HV, V) if pad: o = o[:, :T] return o.to(dtype), (S if output_final_state else None) # =========================================================================== # portable modules # =========================================================================== class ShortConvolution(nn.Conv1d): """Causal depthwise conv1d with an optional silu. Subclasses nn.Conv1d exactly as the fork's does, so the parameter names match and a checkpoint is portable between this path and the Triton one. """ def __init__(self, hidden_size, kernel_size=4, bias=False, activation="silu"): super().__init__(hidden_size, hidden_size, kernel_size, groups=hidden_size, bias=bias) if activation not in (None, "silu", "swish"): raise ValueError(f"unsupported activation {activation!r}") self.hidden_size, self.activation = hidden_size, activation def forward(self, x, cache=None, output_final_state=False, cu_seqlens=None, **kwargs): if cu_seqlens is not None: raise NotImplementedError( "variable-length batching (cu_seqlens) needs the ComplexKDA fla fork") B, T, D = x.shape w = self.kernel_size[0] h = x.transpose(1, 2) if cache is not None: h = torch.cat([cache, h], dim=-1)[:, :, -(T + w - 1):] pad = w - 1 - (h.shape[-1] - T) if pad > 0: h = F.pad(h, (pad, 0)) else: h = F.pad(h, (w - 1, 0)) new_cache = h[:, :, -(w - 1):].contiguous() if output_final_state else None y = self._conv_forward(h, self.weight, self.bias)[:, :, :T].transpose(1, 2) if self.activation in ("silu", "swish"): y = F.silu(y) return y, new_cache class RMSNorm(nn.Module): """rms(x) * weight, with the fork's optional fused residual add. `forward(x, residual, prenorm=True)` returns `(norm(x + residual), x + residual)`. The add is done in the input dtype, matching the fused kernel called with `residual_in_fp32=False`. """ def __init__(self, hidden_size: int, eps: float = 1e-5, elementwise_affine: bool = True): super().__init__() self.hidden_size, self.eps, self.elementwise_affine = hidden_size, eps, elementwise_affine self.weight = nn.Parameter(torch.ones(hidden_size)) if elementwise_affine else None def reset_parameters(self): if self.weight is not None: nn.init.ones_(self.weight) def extra_repr(self) -> str: return f"{self.hidden_size}, eps={self.eps}" def forward(self, x, residual=None, prenorm: bool = False, residual_in_fp32: bool = False): if residual is not None: x = x + (residual.float() if residual_in_fp32 else residual) dt = x.dtype xf = x.float() y = xf * torch.rsqrt(xf.pow(2).mean(-1, keepdim=True) + self.eps) if self.weight is not None: y = y * self.weight.float() y = y.to(dt) return (y, x) if prenorm else y class FusedRMSNormGated(nn.Module): """rms(x) * weight * act(g). The gate is applied AFTER normalising, which is what the fused kernel does and is not interchangeable with gating first.""" def __init__(self, hidden_size, elementwise_affine=True, eps=1e-5, activation="swish"): super().__init__() if activation not in ("swish", "silu", "sigmoid"): raise ValueError(f"Unsupported activation: {activation}") self.hidden_size, self.eps, self.activation = hidden_size, eps, activation self.weight = nn.Parameter(torch.ones(hidden_size)) if elementwise_affine else None def reset_parameters(self): if self.weight is not None: nn.init.ones_(self.weight) def extra_repr(self) -> str: return f"{self.hidden_size}, eps={self.eps}, activation={self.activation}" def forward(self, x, g, **kwargs): dt = x.dtype xf = x.float() y = xf * torch.rsqrt(xf.pow(2).mean(-1, keepdim=True) + self.eps) if self.weight is not None: y = y * self.weight.float() gf = g.float() y = y * (torch.sigmoid(gf) if self.activation == "sigmoid" else gf * torch.sigmoid(gf)) return y.to(dt) class GatedMLP(nn.Module): """SwiGLU: down_proj(swish(gate_proj(x)) * up_proj(x)).""" def __init__(self, hidden_size: int, hidden_ratio: int | None = None, intermediate_size: int | None = None, hidden_act: str = "swish", **kwargs): super().__init__() if hidden_ratio is None: hidden_ratio = 4 if intermediate_size is None: intermediate_size = int(hidden_size * hidden_ratio * 2 / 3) intermediate_size = 256 * ((intermediate_size + 256 - 1) // 256) if hidden_act not in ("swish", "silu"): raise ValueError(f"Unsupported hidden_act: {hidden_act}") self.hidden_size, self.hidden_ratio = hidden_size, hidden_ratio self.intermediate_size, self.hidden_act = intermediate_size, hidden_act 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, x, **kwargs): return self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x)) def _rotate_half(x): x1, x2 = x.chunk(2, dim=-1) return torch.cat((-x2, x1), dim=-1) class RotaryEmbedding(nn.Module): """Rotary position embedding, half-split convention, applied in fp32. Only reached by `use_rope=True` configs; every hybrid published here is NoPE, because the linear layers already carry position. """ def __init__(self, dim: int, base: float = 10000.0): super().__init__() self.dim, self.base = dim, base inv_freq = 1.0 / (base ** (torch.arange(0, dim, 2, dtype=torch.float32) / dim)) self.register_buffer("inv_freq", inv_freq, persistent=False) def forward(self, q, k, seqlen_offset=0, max_seqlen=None, cu_seqlens=None): if cu_seqlens is not None: raise NotImplementedError("variable-length rotary needs the ComplexKDA fla fork") T = q.shape[1] if torch.is_tensor(seqlen_offset): pos = seqlen_offset.view(-1, 1) + torch.arange(T, device=q.device) else: pos = (torch.arange(T, device=q.device) + int(seqlen_offset)).unsqueeze(0) freqs = pos.float().unsqueeze(-1) * self.inv_freq.to(q.device) emb = torch.cat((freqs, freqs), dim=-1) # [B or 1, T, dim] cos, sin = emb.cos().unsqueeze(-2), emb.sin().unsqueeze(-2) qf, kf = q.float(), k.float() q = (qf * cos + _rotate_half(qf) * sin).to(q.dtype) k = (kf * cos + _rotate_half(kf) * sin).to(k.dtype) return q, k class Attention(nn.Module): """The attention layers of a hybrid arm. Causal, through torch SDPA. Two options are not Llama's and both are on in the published hybrids: `output_gate` is Qwen3-Next's sigmoid computed from the LAYER INPUT and applied before `o_proj` (gating after it would scale the residual contribution instead of the per-head mixture, and `o_proj` mixes heads, so the two differ), and `use_rope=False` is NoPE -- no rotary at all, as in Kimi's hybrid. """ def __init__(self, hidden_size: int = 2048, num_heads: int = 32, num_kv_heads: int | None = None, qkv_bias: bool = False, qk_norm: bool = False, output_gate: bool = False, use_rope: bool = True, window_size: int | None = None, rope_theta: float | None = 10000.0, max_position_embeddings: int | None = None, layer_idx: int | None = None): super().__init__() self.hidden_size = hidden_size self.num_heads = num_heads self.num_kv_heads = num_heads if num_kv_heads is None else num_kv_heads self.num_kv_groups = num_heads // self.num_kv_heads self.head_dim = hidden_size // num_heads self.kv_dim = self.num_kv_heads * self.head_dim self.qkv_bias, self.qk_norm, self.output_gate, self.use_rope = qkv_bias, qk_norm, output_gate, use_rope self.window_size, self.rope_theta = window_size, rope_theta self.max_position_embeddings, self.layer_idx = max_position_embeddings, layer_idx self.q_proj = nn.Linear(hidden_size, hidden_size, bias=qkv_bias) self.k_proj = nn.Linear(hidden_size, self.kv_dim, bias=qkv_bias) self.v_proj = nn.Linear(hidden_size, self.kv_dim, bias=qkv_bias) self.o_proj = nn.Linear(hidden_size, hidden_size, bias=False) if output_gate: self.g_proj = nn.Linear(hidden_size, hidden_size, bias=False) if qk_norm: self.q_norm = RMSNorm(self.head_dim) self.k_norm = RMSNorm(self.head_dim) self.rotary = RotaryEmbedding(dim=self.head_dim, base=self.rope_theta) if use_rope else None def forward(self, hidden_states, attention_mask=None, past_key_values=None, output_attentions: bool = False, use_cache: bool = False, **kwargs): if attention_mask is not None and attention_mask.dim() != 2: raise ValueError( "Expected attention_mask as a 0-1 matrix of shape [batch_size, seq_len] " "(0 = padding). Arbitrary [b, q, k] masks are not supported.") if kwargs.get("cu_seqlens") is not None: raise NotImplementedError("variable-length attention needs the ComplexKDA fla fork") B, q_len, _ = hidden_states.shape q = self.q_proj(hidden_states).view(B, q_len, self.num_heads, self.head_dim) k = self.k_proj(hidden_states).view(B, q_len, self.num_kv_heads, self.head_dim) v = self.v_proj(hidden_states).view(B, q_len, self.num_kv_heads, self.head_dim) if self.qk_norm: q, k = self.q_norm(q), self.k_norm(k) seqlen_offset = 0 if past_key_values is not None: seqlen_offset = past_key_values.get_seq_length(self.layer_idx) if attention_mask is not None: # Padding sits on the LEFT of a padded batch, so a row's real # position is its offset minus its padding. lens = attention_mask.sum(-1, dtype=torch.long) seqlen_offset = seqlen_offset + lens - attention_mask.shape[-1] if self.rotary is not None: q, k = self.rotary(q, k, seqlen_offset=seqlen_offset) if past_key_values is not None: k, v = past_key_values.update_attn(self.layer_idx, k, v, window_size=self.window_size) # [B, T, H, D] -> [B, H, T, D] qt, kt, vt = (x.transpose(1, 2) for x in (q, k, v)) k_len = kt.shape[2] attn_bias = None is_causal = False if q_len == k_len and attention_mask is None and self.window_size is None: is_causal = True else: pos_q = torch.arange(k_len - q_len, k_len, device=q.device) pos_k = torch.arange(k_len, device=q.device) # Bottom-right alignment: query t attends keys <= its own position. keep = pos_k[None, :] <= pos_q[:, None] if self.window_size is not None: keep &= pos_k[None, :] > pos_q[:, None] - self.window_size keep = keep[None, None] if attention_mask is not None: pad = attention_mask[:, None, None, :].bool() if pad.shape[-1] != k_len: pad = F.pad(pad, (k_len - pad.shape[-1], 0), value=True) keep = keep & pad attn_bias = torch.zeros(keep.shape, dtype=qt.dtype, device=q.device) attn_bias = attn_bias.masked_fill(~keep, torch.finfo(qt.dtype).min) gqa = {"enable_gqa": True} if self.num_kv_groups > 1 else {} o = F.scaled_dot_product_attention(qt, kt, vt, attn_mask=attn_bias, is_causal=is_causal, **gqa) o = o.transpose(1, 2).reshape(B, q_len, -1) if self.output_gate: o = o * torch.sigmoid(self.g_proj(hidden_states)) return self.o_proj(o), None, past_key_values # =========================================================================== # the mixer # =========================================================================== def _identity(x): return x class ComplexKimiDeltaAttention(nn.Module): """Kimi Delta Attention whose decay gate may be negative. Beyond KimiDeltaAttention: * `gate` selects the decay parameterisation (two unsigned, two signed); * the `+-1` part of a signed gate is carried as a running sign pushed onto q/k (the gauge), so the recurrence itself still runs on `|alpha|`. On the Triton path the sign is handed to the op as `sign=` and applied inside KDA's own l2norm epilogue; the torch path gauges explicitly here. """ def __init__(self, hidden_size: int = 2048, expand_v: float = 1, head_dim: int = 128, num_heads: int = 16, num_v_heads: int | None = None, mode: str = "chunk", use_short_conv: bool = True, allow_neg_eigval: bool = False, gate: str = "signed_sigmoid2", drop_silu: bool = False, drop_key_silu: bool = False, conv_silu: str = "qkv", gate_init_style: str = "shipped", output_gate: str = "lowrank", beta_init_style: str = "standard", lower_bound: float = -5.0, conv_size: int = 4, conv_bias: bool = False, layer_idx: int | None = None, norm_eps: float = 1e-5, chunk_size: int | None = None, **kwargs): super().__init__() if gate not in GATES: raise ValueError(f"gate must be one of {GATES}, got {gate!r}") if gate_init_style not in ("shipped", "spread"): raise ValueError(f"gate_init_style must be 'shipped' or 'spread', got {gate_init_style!r}") if output_gate not in ("lowrank", "linear"): raise ValueError(f"output_gate must be 'lowrank' or 'linear', got {output_gate!r}") if beta_init_style not in ("standard", "spread"): raise ValueError(f"beta_init_style must be 'standard' or 'spread', got {beta_init_style!r}") if mode not in ("chunk", "fused_recurrent"): raise ValueError(f"unsupported mode {mode!r}") if not (-5 <= lower_bound < 0): raise ValueError(f"lower_bound must be in [-5, 0), got {lower_bound}") self.mode = mode self.allow_neg_eigval = allow_neg_eigval self.gate = gate self.act = _identity if drop_silu else F.silu self.k_act = _identity if drop_silu or drop_key_silu else F.silu self.gate_init_style = gate_init_style self.beta_init_style = beta_init_style self.safe_gate = safe_gate_ok(gate) self.lower_bound = lower_bound self.hidden_size = hidden_size self.expand_v = expand_v self.chunk_size = CHUNK_SIZE if chunk_size is None else chunk_size self.use_short_conv = use_short_conv self.conv_size = conv_size self.conv_bias = conv_bias self.head_dim = head_dim self.num_heads = num_heads self.num_v_heads = num_heads if num_v_heads is None else num_v_heads self.head_k_dim = head_dim self.head_v_dim = int(head_dim * expand_v) self.key_dim = int(self.num_heads * self.head_k_dim) self.value_dim = int(self.num_v_heads * self.head_v_dim) self.layer_idx = layer_idx if not math.isclose(head_dim * expand_v, self.head_v_dim, rel_tol=1e-5): raise ValueError(f"expand_v={expand_v} does not give an integer head_v_dim from head_dim={head_dim}") if self.num_v_heads > self.num_heads and self.num_v_heads % self.num_heads != 0: raise ValueError(f"num_v_heads={self.num_v_heads} must be divisible by num_heads={self.num_heads}") if self.num_v_heads > self.num_heads and is_signed(gate): warnings.warn( "signed gate under GVA expands q/k to num_v_heads (the gauge is per value head " "but q/k are shared), losing the GVA memory saving.", stacklevel=2) self.q_proj = nn.Linear(hidden_size, self.key_dim, bias=False) self.k_proj = nn.Linear(hidden_size, self.key_dim, bias=False) self.v_proj = nn.Linear(hidden_size, self.value_dim, bias=False) if any(c not in "qkv" for c in conv_silu): raise ValueError(f"conv_silu must be a subset of 'qkv', got {conv_silu!r}") self.conv_silu = "" if drop_silu else conv_silu if drop_key_silu: self.conv_silu = self.conv_silu.replace("k", "") conv_cls = _FLA_SHORTCONV or ShortConvolution if use_short_conv: self.q_conv1d = conv_cls(hidden_size=self.key_dim, kernel_size=conv_size, bias=conv_bias, activation="silu" if "q" in self.conv_silu else None) self.k_conv1d = conv_cls(hidden_size=self.key_dim, kernel_size=conv_size, bias=conv_bias, activation="silu" if "k" in self.conv_silu else None) self.v_conv1d = conv_cls(hidden_size=self.value_dim, kernel_size=conv_size, bias=conv_bias, activation="silu" if "v" in self.conv_silu else None) self.gate_dim = int(self.num_v_heads * self.head_k_dim) self.f_proj = nn.Sequential( nn.Linear(hidden_size, self.head_v_dim, bias=False), nn.Linear(self.head_v_dim, self.gate_dim, bias=False), ) self.b_proj = nn.Linear(hidden_size, self.num_v_heads, bias=beta_init_style == "spread") if self.safe_gate: self.A_log = nn.Parameter(torch.zeros(self.num_v_heads, dtype=torch.float32)) else: self.A_log = nn.Parameter(torch.log(torch.empty(self.num_v_heads, dtype=torch.float32).uniform_(1, 16))) self.A_log._no_weight_decay = True self.dt_bias = nn.Parameter(init_dt_bias(gate, self.gate_dim, lower_bound, gate_init_style)) self.dt_bias._no_weight_decay = True # The output forget gate. "lowrank" is fla's and Kimi Linear's factored # hidden -> head_v_dim -> value_dim pair; "linear" is Kimi K3's single # full-rank map. Both sit downstream of the recurrence, so neither # interacts with the signed decay gate. if output_gate == "lowrank": self.g_proj = nn.Sequential( nn.Linear(hidden_size, self.head_v_dim, bias=False), nn.Linear(self.head_v_dim, self.value_dim, bias=True), ) else: self.g_proj = nn.Linear(hidden_size, self.value_dim, bias=True) norm_gated_cls = _FLA_RMSNORM_GATED or FusedRMSNormGated self.o_norm = norm_gated_cls(self.head_v_dim, activation="sigmoid", eps=norm_eps) self.o_proj = nn.Linear(self.value_dim, hidden_size, bias=False) def forward(self, hidden_states, attention_mask=None, past_key_values=None, use_cache: bool | None = False, output_attentions: bool | None = False, **kwargs): if attention_mask is not None and attention_mask.dim() != 2: raise ValueError( "Expected attention_mask as a 0-1 matrix of shape [batch_size, seq_len] " "(0 = padding). Arbitrary [b, q, k] masks are not supported.") if kwargs.get("cu_seqlens") is not None and not USE_KERNEL: raise NotImplementedError("variable-length batching needs the ComplexKDA fla fork") cu_seqlens = kwargs.get("cu_seqlens") B, q_len, _ = hidden_states.shape last_state = None if past_key_values is not None and self.layer_idx is not None: last_state = past_key_values.get(self.layer_idx) if self.use_short_conv: cq, ck, cv = last_state["conv_state"] if last_state is not None else (None, None, None) q, cq = self.q_conv1d(x=self.q_proj(hidden_states), cache=cq, output_final_state=use_cache, cu_seqlens=cu_seqlens) k, ck = self.k_conv1d(x=self.k_proj(hidden_states), cache=ck, output_final_state=use_cache, cu_seqlens=cu_seqlens) v, cv = self.v_conv1d(x=self.v_proj(hidden_states), cache=cv, output_final_state=use_cache, cu_seqlens=cu_seqlens) else: cq = ck = cv = None q = self.act(self.q_proj(hidden_states)) k = self.k_act(self.k_proj(hidden_states)) v = self.act(self.v_proj(hidden_states)) g = self.f_proj(hidden_states) beta = self.b_proj(hidden_states) q = q.view(*q.shape[:-1], -1, self.head_k_dim) k = k.view(*k.shape[:-1], -1, self.head_k_dim) g = g.view(*g.shape[:-1], -1, self.head_k_dim) v = v.view(*v.shape[:-1], -1, self.head_v_dim) sign, g = compute_gate(self.gate, g, self.A_log, self.dt_bias, self.lower_bound) # GVA: the gauge is per value head but q/k are shared across the group, # so q/k are expanded to HV before either path applies it. if sign is not None and sign.shape[2] != q.shape[2]: r = sign.shape[2] // q.shape[2] q, k = q.repeat_interleave(r, dim=2), k.repeat_interleave(r, dim=2) recurrent_state = last_state["recurrent_state"] if last_state is not None else None scale = self.head_k_dim ** -0.5 P = None if USE_KERNEL: # The kernels take `sign` directly: the gauge rides KDA's own l2norm # epilogue, and the final state comes back already un-gauged. mode = "fused_recurrent" if (q_len <= 64 and not self.training) else self.mode op = _CHUNK_KDA if mode == "chunk" else _FUSED_RECURRENT_KDA extra = dict(use_gate_in_kernel=False, safe_gate=self.safe_gate) if mode == "chunk" else {} o, recurrent_state = op( q=q, k=k, v=v, g=g, beta=beta, sign=sign, scale=scale, initial_state=recurrent_state, output_final_state=bool(use_cache), use_qk_l2norm_in_kernel=True, use_beta_sigmoid_in_kernel=True, allow_neg_eigval=self.allow_neg_eigval, lower_bound=self.lower_bound, state_v_first=True, cu_seqlens=cu_seqlens, **extra) state_v_first = True else: if sign is not None: P = running_sign(sign, cu_seqlens) q, k = apply_sign(q, P), apply_sign(k, P) qn = F.normalize(q.float(), dim=-1, eps=1e-6).to(q.dtype) kn = F.normalize(k.float(), dim=-1, eps=1e-6).to(k.dtype) bt = torch.sigmoid(beta.float()) * (2.0 if self.allow_neg_eigval else 1.0) fn = recurrent_kda_torch if q_len <= 8 else chunk_kda_torch extra = {} if fn is recurrent_kda_torch else dict(chunk_size=min(self.chunk_size, max(q_len, 1))) o, recurrent_state = fn( qn, kn, v, g.to(q.dtype), bt.to(q.dtype), scale=scale, initial_state=recurrent_state, output_final_state=bool(use_cache), **extra) state_v_first = False if P is not None and recurrent_state is not None: recurrent_state = ungauge_state(recurrent_state, P[:, -1], state_v_first=state_v_first, head_k_dim=self.head_k_dim) if use_cache and past_key_values is not None and self.layer_idx is not None: past_key_values.update_recurrent( self.layer_idx, recurrent_state=recurrent_state, conv_state=(cq, ck, cv) if self.use_short_conv else None, state_v_first=state_v_first, offset=q_len, ) g_out = self.g_proj(hidden_states) o = self.o_norm(o, g_out.view(*g_out.shape[:-1], -1, self.head_v_dim)) o = o.reshape(*o.shape[:-2], -1) return self.o_proj(o), None, past_key_values # =========================================================================== # cache # =========================================================================== class ComplexKDACache: """Per-layer state for incremental decoding. Not a `transformers.Cache`: that class models a growing key/value pair per layer, and a linear-attention layer has a FIXED-SIZE recurrent state plus a short-convolution window instead. Hybrid arms hold both kinds, which is why the two kinds of entry live side by side here. A state is stored in whichever layout produced it (`state_v_first` records which), so a cache filled by the Triton path and one filled by the torch path are not interchangeable -- the flag makes that a loud error instead of a transposed state. """ # Attributes transformers' generation loop probes on whatever cache it was # handed. They are plain class attributes rather than properties so that a # version which reads one this does not define fails on the name it wants # rather than on something further downstream. is_compileable = False is_sliding = False def __init__(self, seen_tokens: int = 0): self.states: dict[int, dict[str, Any]] = {} self._seen_tokens = seen_tokens def __len__(self) -> int: return len(self.states) def get(self, layer_idx: int): return self.states.get(layer_idx) def get_seq_length(self, layer_idx: int = 0) -> int: state = self.states.get(layer_idx) return 0 if state is None else state.get("offset", 0) def get_max_cache_shape(self, layer_idx: int = 0) -> int | None: return None def update_recurrent(self, layer_idx: int, recurrent_state, conv_state, state_v_first: bool, offset: int): prev = self.states.get(layer_idx) if prev is not None and prev.get("state_v_first") != state_v_first: raise ValueError( f"layer {layer_idx}: cached state was written with state_v_first=" f"{prev.get('state_v_first')} and is being updated with {state_v_first}. " "The kernel and torch backends store the state on opposite axes; do not " "switch COMPLEX_KDA_BACKEND part-way through a generation.") self.states[layer_idx] = { "recurrent_state": recurrent_state, "conv_state": conv_state, "state_v_first": state_v_first, "offset": self.get_seq_length(layer_idx) + offset, } def update_attn(self, layer_idx: int, k: torch.Tensor, v: torch.Tensor, window_size: int | None = None): """Append these keys/values and return the full history. The offset counts TOKENS SEEN, not calls: a prefill hands over many at once, and under a sliding window it keeps counting after the cache has stopped growing. It is what `Attention` rotates by, so getting it from `k.shape[1]` would put a windowed model's rotary back at the start of the window on every step. """ prev = self.states.get(layer_idx) n_new = k.shape[1] if prev is not None and prev.get("attn_state") is not None: pk, pv = prev["attn_state"] k, v = torch.cat([pk, k], dim=1), torch.cat([pv, v], dim=1) # TRIM WHAT IS STORED, RETURN THE WHOLE CONCATENATION. Trimming before # the caller attends would hand a prefill of T > window only the last # `window` keys for ALL T queries -- the early ones would then attend a # window that starts after them. The caller applies the window mask; # this only bounds what the NEXT step has to carry, and since what was # stored is already within the window, the concatenation returned on a # decode step is at most `window + 1` long. stored = (k, v) if window_size is None else (k[:, -window_size:], v[:, -window_size:]) self.states[layer_idx] = { "attn_state": stored, "offset": (0 if prev is None else prev.get("offset", 0)) + n_new, } return k, v def reorder_cache(self, beam_idx: torch.LongTensor): for state in self.states.values(): for key in ("recurrent_state",): if state.get(key) is not None: state[key] = state[key].index_select(0, beam_idx.to(state[key].device)) if state.get("conv_state") is not None: state["conv_state"] = tuple( None if c is None else c.index_select(0, beam_idx.to(c.device)) for c in state["conv_state"]) if state.get("attn_state") is not None: state["attn_state"] = tuple( t.index_select(0, beam_idx.to(t.device)) for t in state["attn_state"]) return self # =========================================================================== # the model # =========================================================================== class ComplexKDABlock(nn.Module): def __init__(self, config: ComplexKDAConfig, layer_idx: int): super().__init__() self.config = config self.layer_idx = layer_idx norm_cls = _FLA_RMSNORM or RMSNorm self.attn_norm = norm_cls(config.hidden_size, eps=config.norm_eps) spec = get_hybrid_attention_spec(config.attn, layer_idx=layer_idx) if spec is not None: # `qk_norm`, `output_gate` and `use_rope` are read with .get: they # are optional keys that the config preserves rather than fields of # the spec, and their defaults are Attention's own. The published # hybrids need the last two -- their attention is GATED and NoPE -- # and without them an exported hybrid is a different model: rotary # where the run had none, and no `g_proj` at all. self.attn = Attention( hidden_size=config.hidden_size, num_heads=spec["num_heads"], num_kv_heads=spec["num_kv_heads"], qkv_bias=spec["qkv_bias"], qk_norm=spec.get("qk_norm", False), output_gate=spec.get("output_gate", False), use_rope=spec.get("use_rope", True), window_size=spec["window_size"], rope_theta=spec["rope_theta"], max_position_embeddings=config.max_position_embeddings, layer_idx=layer_idx, ) else: self.attn = ComplexKimiDeltaAttention( mode=config.attn_mode, hidden_size=config.hidden_size, expand_v=config.expand_v, head_dim=config.head_dim, num_heads=config.num_heads, num_v_heads=config.num_v_heads, use_short_conv=config.use_short_conv, drop_silu=config.drop_silu, drop_key_silu=config.drop_key_silu, allow_neg_eigval=config.allow_neg_eigval, gate=config.gate, gate_init_style=config.gate_init_style, output_gate=config.output_gate, conv_silu=config.conv_silu, beta_init_style=config.beta_init_style, lower_bound=config.lower_bound, conv_size=config.conv_size, norm_eps=config.norm_eps, layer_idx=layer_idx, ) self.mlp_norm = norm_cls(config.hidden_size, eps=config.norm_eps) self.mlp = GatedMLP( hidden_size=config.hidden_size, hidden_ratio=config.hidden_ratio, intermediate_size=config.intermediate_size, hidden_act=config.hidden_act, ) def forward(self, hidden_states, attention_mask=None, past_key_values=None, use_cache: bool | None = False, output_attentions: bool | None = False, **kwargs): residual = hidden_states hidden_states = self.attn_norm(hidden_states) hidden_states, attentions, past_key_values = self.attn( hidden_states=hidden_states, attention_mask=attention_mask, past_key_values=past_key_values, use_cache=use_cache, output_attentions=output_attentions, **kwargs, ) hidden_states, residual = self.mlp_norm(hidden_states, residual, True) hidden_states = self.mlp(hidden_states) return residual + hidden_states, attentions, past_key_values class ComplexKDAPreTrainedModel(PreTrainedModel): config_class = ComplexKDAConfig base_model_prefix = "model" supports_gradient_checkpointing = True _no_split_modules = ["ComplexKDABlock"] _supports_sdpa = True _can_compile_fullgraph = False def _init_weights(self, module: nn.Module): std = self.config.initializer_range if isinstance(module, ComplexKimiDeltaAttention): if next(module.parameters()).device.type != "meta": with torch.no_grad(): module.A_log.zero_() dt = torch.exp( torch.rand_like(module.dt_bias) * (math.log(0.1) - math.log(0.001)) + math.log(0.001) ).clamp(min=1e-4) module.dt_bias.copy_(init_dt_bias( module.gate, lower_bound=module.lower_bound, gate_init_style=module.gate_init_style, dt=dt)) return if isinstance(module, (nn.Linear, nn.Conv1d)): nn.init.normal_(module.weight, mean=0.0, std=std) if module.bias is not None: nn.init.zeros_(module.bias) elif isinstance(module, nn.Embedding): nn.init.normal_(module.weight, mean=0.0, std=std) elif hasattr(module, "reset_parameters"): module.reset_parameters() class ComplexKDAModel(ComplexKDAPreTrainedModel): def __init__(self, config: ComplexKDAConfig): super().__init__(config) self.padding_idx = config.pad_token_id self.vocab_size = config.vocab_size self.embeddings = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx) self.layers = nn.ModuleList( [ComplexKDABlock(config, i) for i in range(config.num_hidden_layers)]) self.norm = (_FLA_RMSNORM or RMSNorm)(config.hidden_size, eps=config.norm_eps) self.gradient_checkpointing = False self.post_init() def get_input_embeddings(self): return self.embeddings def set_input_embeddings(self, value): self.embeddings = value def forward(self, input_ids=None, attention_mask=None, inputs_embeds=None, past_key_values=None, use_cache=None, output_attentions=None, output_hidden_states=None, return_dict=None, **kwargs): if output_attentions: warnings.warn("ComplexKDAModel does not support `output_attentions`; setting it to False.") output_attentions = False output_hidden_states = (output_hidden_states if output_hidden_states is not None else self.config.output_hidden_states) use_cache = use_cache if use_cache is not None else (self.config.use_cache and not self.training) return_dict = return_dict if return_dict is not None else True if input_ids is not None and inputs_embeds is not None: raise ValueError("You cannot specify both input_ids and inputs_embeds at the same time") if input_ids is None and inputs_embeds is None: raise ValueError("You have to specify either input_ids or inputs_embeds") hidden_states = self.embeddings(input_ids) if inputs_embeds is None else inputs_embeds if use_cache and past_key_values is None: past_key_values = ComplexKDACache() if past_key_values is not None and not isinstance(past_key_values, ComplexKDACache): raise TypeError( f"ComplexKDA needs a ComplexKDACache (it stores recurrent state, not key/value " f"pairs); got {type(past_key_values).__name__}.") all_hidden_states = () if output_hidden_states else None for layer in self.layers: if output_hidden_states: all_hidden_states += (hidden_states,) if self.gradient_checkpointing and self.training: hidden_states, _, past_key_values = self._gradient_checkpointing_func( layer.__call__, hidden_states, attention_mask, past_key_values, use_cache, output_attentions, **kwargs) else: hidden_states, _, past_key_values = layer( hidden_states, attention_mask=attention_mask, past_key_values=past_key_values, use_cache=use_cache, output_attentions=output_attentions, **kwargs) hidden_states = self.norm(hidden_states) if output_hidden_states: all_hidden_states += (hidden_states,) if not return_dict: return tuple(x for x in (hidden_states, past_key_values, all_hidden_states) if x is not None) return BaseModelOutputWithPast( last_hidden_state=hidden_states, past_key_values=past_key_values, hidden_states=all_hidden_states, attentions=None, ) def _tied_weights_keys_declaration(): """How this transformers spells "lm_head.weight IS the embedding". Every ladder cell ties its embeddings (the 1.3B arms do not), and the two transformers generations declare that differently: 4.x a LIST of regex patterns matched against parameter names 5.x a {target: source} MAPPING THE WRONG ONE IS NOT A WARNING. A list under transformers 5 raises `'list' object has no attribute 'keys'` from inside `post_init` -- for every tied checkpoint, and only for tied ones, so it passes every test run against an untied model and then fails for most of the release. """ mapping = {"lm_head.weight": "model.embeddings.weight"} patterns = ["lm_head.weight"] try: from transformers.modeling_utils import PreTrainedModel annotation = str(getattr(PreTrainedModel, "__annotations__", {}) .get("_tied_weights_keys", "")) if annotation: return mapping if "dict" in annotation.lower() else patterns except Exception: pass try: import transformers return mapping if int(str(transformers.__version__).split(".")[0]) >= 5 else patterns except Exception: return patterns class ComplexKDAForCausalLM(ComplexKDAPreTrainedModel, GenerationMixin): _tied_weights_keys = _tied_weights_keys_declaration() def __init__(self, config: ComplexKDAConfig): super().__init__(config) self.model = ComplexKDAModel(config) self.vocab_size = config.vocab_size self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) self.criterion = None self.post_init() def get_input_embeddings(self): return self.model.embeddings def set_input_embeddings(self, value): self.model.embeddings = value def get_output_embeddings(self): return self.lm_head def set_output_embeddings(self, new_embeddings): self.lm_head = new_embeddings def get_decoder(self): return self.model def set_decoder(self, decoder): self.model = decoder def tie_weights(self, *args, **kwargs): """Tie the head to the embedding OURSELVES, rather than describing the tie and hoping this transformers acts on the description. Every ladder cell ties, and the exporters write NO `lm_head.weight` for a tied geometry -- there is no second tensor to write. transformers 5.3 nonetheless decides the key "is present in the checkpoint", declines to tie, and leaves the head on the META device: `from_pretrained` returns without error and the first `.to(device)` raises "Cannot copy out of meta tensor". A CPU-only smoke test does not even get that far -- it returns a model whose head is data-less. Doing the assignment here is version-independent, and transformers calls this both in `post_init` and after loading the weights, so the alias survives materialisation. """ if getattr(self.config, "tie_word_embeddings", False): embeddings = self.get_input_embeddings() if embeddings is not None: self.lm_head.weight = embeddings.weight # *args/**kwargs: transformers 5.6 passes `recompute_mapping`, 4.x # passes nothing. Forward whatever it sends rather than pinning a # signature that one of them will not call. return super().tie_weights(*args, **kwargs) def forward(self, input_ids=None, attention_mask=None, inputs_embeds=None, past_key_values=None, labels=None, use_cache=None, output_attentions=None, output_hidden_states=None, return_dict=None, logits_to_keep=0, **kwargs): return_dict = return_dict if return_dict is not None else True outputs = self.model( input_ids=input_ids, attention_mask=attention_mask, inputs_embeds=inputs_embeds, past_key_values=past_key_values, use_cache=use_cache, output_attentions=output_attentions, output_hidden_states=output_hidden_states, return_dict=return_dict, **kwargs) hidden_states = outputs[0] if logits_to_keep: hidden_states = hidden_states[:, -logits_to_keep:] logits = self.lm_head(hidden_states) loss = None if labels is not None: criterion = self.criterion if self.criterion is not None else nn.CrossEntropyLoss() labels = labels.to(logits.device) # Shift here rather than on the logits, matching the training stack. labels = torch.cat((labels[..., 1:], torch.full_like(labels[:, :1], criterion.ignore_index)), 1) loss = criterion(logits.view(labels.numel(), -1), labels.view(-1)) if not return_dict: output = (logits,) + tuple(outputs[1:]) return (loss,) + output if loss is not None else output return CausalLMOutputWithPast( loss=loss, logits=logits, past_key_values=outputs.past_key_values, hidden_states=outputs.hidden_states, attentions=None, ) # ---- generation ------------------------------------------------------- # # `generate` builds a DynamicCache by default, which this model cannot use # (see ComplexKDACache). Installing ours here is the documented escape # route: a cache already present in model_kwargs is left alone. def _prepare_cache_for_generation(self, generation_config, model_kwargs, *args, **kwargs): # *args absorbs the positional tail, which differs across transformers # versions; the two arguments this needs have not moved. if generation_config.use_cache and model_kwargs.get("past_key_values") is None: model_kwargs["past_key_values"] = ComplexKDACache() return True return False def prepare_inputs_for_generation(self, input_ids, past_key_values=None, attention_mask=None, inputs_embeds=None, use_cache=True, logits_to_keep=None, **kwargs): if past_key_values is not None and len(past_key_values) > 0: input_ids = input_ids[:, -1:] model_inputs = {"input_ids": input_ids, "inputs_embeds": None} if inputs_embeds is not None and past_key_values is None: model_inputs = {"input_ids": None, "inputs_embeds": inputs_embeds} model_inputs.update( past_key_values=past_key_values, use_cache=use_cache, attention_mask=attention_mask, logits_to_keep=1 if logits_to_keep is None else logits_to_keep, ) return model_inputs def _reorder_cache(self, past_key_values, beam_idx): return past_key_values.reorder_cache(beam_idx)