kda-sigmoid-hybrid-1.3B-100B / modeling_complex_kda.py
korbip's picture
Point the fla fork at OpenEuroLLM/ComplexKDA
1b9ccd8 verified
Raw History Blame Contribute Delete
63 kB
# 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 * <k_c * alpha(i,c], k_i> 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)