Text Generation
Transformers
Safetensors
complex_kda
complex-kda
linear-attention
kimi-delta-attention
conversational
custom_code
Instructions to use openeurollm/kda-sigmoid-hybrid-1.3B-100B with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Transformers
How to use openeurollm/kda-sigmoid-hybrid-1.3B-100B with Transformers:
# Use a pipeline as a high-level helper from transformers import pipeline pipe = pipeline("text-generation", model="openeurollm/kda-sigmoid-hybrid-1.3B-100B", trust_remote_code=True) messages = [ {"role": "user", "content": "Who are you?"}, ] pipe(messages)# Load model directly from transformers import AutoModelForCausalLM model = AutoModelForCausalLM.from_pretrained("openeurollm/kda-sigmoid-hybrid-1.3B-100B", trust_remote_code=True, device_map="auto") - Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- vLLM
How to use openeurollm/kda-sigmoid-hybrid-1.3B-100B with vLLM:
Install from pip and serve model
# Install vLLM from pip: pip install vllm # Start the vLLM server: vllm serve "openeurollm/kda-sigmoid-hybrid-1.3B-100B" # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:8000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "openeurollm/kda-sigmoid-hybrid-1.3B-100B", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker
docker model run hf.co/openeurollm/kda-sigmoid-hybrid-1.3B-100B
- SGLang
How to use openeurollm/kda-sigmoid-hybrid-1.3B-100B with SGLang:
Install from pip and serve model
# Install SGLang from pip: pip install sglang # Start the SGLang server: python3 -m sglang.launch_server \ --model-path "openeurollm/kda-sigmoid-hybrid-1.3B-100B" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "openeurollm/kda-sigmoid-hybrid-1.3B-100B", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }'Use Docker images
docker run --gpus all \ --shm-size 32g \ -p 30000:30000 \ -v ~/.cache/huggingface:/root/.cache/huggingface \ --env "HF_TOKEN=<secret>" \ --ipc=host \ lmsysorg/sglang:latest \ python3 -m sglang.launch_server \ --model-path "openeurollm/kda-sigmoid-hybrid-1.3B-100B" \ --host 0.0.0.0 \ --port 30000 # Call the server using curl (OpenAI-compatible API): curl -X POST "http://localhost:30000/v1/chat/completions" \ -H "Content-Type: application/json" \ --data '{ "model": "openeurollm/kda-sigmoid-hybrid-1.3B-100B", "messages": [ { "role": "user", "content": "What is the capital of France?" } ] }' - Docker Model Runner
How to use openeurollm/kda-sigmoid-hybrid-1.3B-100B with Docker Model Runner:
docker model run hf.co/openeurollm/kda-sigmoid-hybrid-1.3B-100B
Download modeling_complex_kda.py from openeurollm/kda-sigmoid-hybrid-1.3B-100B: direct link, hf CLI and curl.
- Browser
- Download file 63 kB
-
https://huggingface.co/openeurollm/kda-sigmoid-hybrid-1.3B-100B/resolve/main/modeling_complex_kda.py
- Command line
-
hf download hf://openeurollm/kda-sigmoid-hybrid-1.3B-100B/modeling_complex_kda.py
-
curl -L -o modeling_complex_kda.py https://huggingface.co/openeurollm/kda-sigmoid-hybrid-1.3B-100B/resolve/main/modeling_complex_kda.py
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.""" | |
| def forward(ctx, x, P): | |
| ctx.save_for_backward(P) | |
| return x * P.to(x.dtype) | |
| 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) | |