Initial K3 snapshot: 0.5B KDA/MLA/MoE train path

Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias
backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
This commit is contained in:
dela
2026-08-25 14:43:17 +08:00
commit 584f7e9e73
140 changed files with 21592 additions and 0 deletions
+5
View File
@@ -0,0 +1,5 @@
"""KDA operator API and implementation backends."""
from .api import chunk_kda
__all__ = ["chunk_kda"]
+167
View File
@@ -0,0 +1,167 @@
"""Training-facing KDA operator with the same boundary as FLA's ``chunk_kda``."""
from __future__ import annotations
import warnings
from functools import lru_cache
import torch
import torch.nn.functional as F
from .reference.chunkwise import DECAY_BLOCK, _EXP_LIMIT, naive_chunk_kda
@lru_cache(maxsize=1)
def _fla_chunk_kda():
try:
from fla.ops.kda import chunk_kda
except ImportError:
return None
return chunk_kda
def _reference_chunk_size(T: int, requested: int) -> int:
size = min(T, requested)
while T % size:
size -= 1
return size
def chunk_kda(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
g: torch.Tensor,
beta: torch.Tensor,
*,
A_log: torch.Tensor | None = None,
dt_bias: torch.Tensor | None = None,
scale: float | None = None,
initial_state: torch.Tensor | None = None,
output_final_state: bool = False,
use_qk_l2norm_in_kernel: bool = False,
use_gate_in_kernel: bool = False,
use_beta_sigmoid_in_kernel: bool = False,
safe_gate: bool = False,
lower_bound: float | None = None,
chunk_size: int = 64,
backend: str = "reference",
):
"""Run an explicitly selected KDA implementation.
``reference`` and its legacy alias ``torch`` use this repository's
differentiable PyTorch implementation. ``triton`` uses the vendored
FLA NVIDIA Triton kernels in ``kda._fla`` (chunk_size 32 or 64,
CUDA). ``fla`` is reserved for explicit upstream parity runs.
"""
supported = {"reference", "triton", "fla", "torch", "auto"}
if backend not in supported:
raise ValueError(f"backend must be one of {sorted(supported)}")
if not use_qk_l2norm_in_kernel:
# Backend-independent: this is a property of the recurrence, not of any
# one implementation.
warnings.warn(
"use_qk_l2norm_in_kernel=False: KDA's chunkwise form assumes "
"||k||=1 so that I + tril(A_kk*beta) has a convergent Neumann "
"series. Unnormalised k makes the exact output grow like "
"||k||^chunk_size and can reach inf on any backend.",
RuntimeWarning,
stacklevel=2,
)
if backend == "auto":
warnings.warn(
"backend='auto' is deprecated and now selects the local reference backend; "
"use backend='fla' explicitly for upstream FLA",
DeprecationWarning,
stacklevel=2,
)
backend = "reference"
if backend == "torch":
backend = "reference"
if backend == "triton":
from .triton.chunk import chunk_kda as triton_chunk_kda
fla_chunk = 32 if chunk_size <= 32 else 64
return triton_chunk_kda(
q,
k,
v,
g,
beta,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
chunk_size=fla_chunk,
use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
use_gate_in_kernel=use_gate_in_kernel,
use_beta_sigmoid_in_kernel=use_beta_sigmoid_in_kernel,
A_log=A_log,
dt_bias=dt_bias,
safe_gate=safe_gate,
lower_bound=lower_bound,
)
if backend == "fla":
fused_op = _fla_chunk_kda()
if fused_op is None:
raise RuntimeError(
"backend='fla' requires a complete flash-linear-attention installation"
)
return fused_op(
q,
k,
v,
g,
beta,
A_log=A_log,
dt_bias=dt_bias,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
use_qk_l2norm_in_kernel=use_qk_l2norm_in_kernel,
use_gate_in_kernel=use_gate_in_kernel,
use_beta_sigmoid_in_kernel=use_beta_sigmoid_in_kernel,
safe_gate=safe_gate,
lower_bound=lower_bound,
chunk_size=32 if chunk_size <= 32 else 64,
)
if safe_gate and lower_bound is not None:
# _decayed_dot exponentiates at most DECAY_BLOCK steps of gate decay,
# and safe_gate bounds each step by |lower_bound|.
budget = DECAY_BLOCK * abs(lower_bound)
if budget > _EXP_LIMIT:
raise ValueError(
f"lower_bound={lower_bound} allows a gate span of {budget:.1f} "
f"per {DECAY_BLOCK}-row block, which overflows exp() "
f"(limit {_EXP_LIMIT:.1f}) and yields NaN. Use "
f"|lower_bound| < {_EXP_LIMIT / DECAY_BLOCK:.2f} or "
"backend='triton'."
)
if use_qk_l2norm_in_kernel:
q, k = F.normalize(q, dim=-1), F.normalize(k, dim=-1)
if use_beta_sigmoid_in_kernel:
beta = beta.sigmoid()
if use_gate_in_kernel:
if A_log is None:
raise ValueError("A_log is required when use_gate_in_kernel=True")
bias = 0 if dt_bias is None else dt_bias.view(g.shape[-2:])
gate_input = g + bias
rate = A_log.exp().view(1, 1, -1, 1)
if safe_gate:
if lower_bound is None:
raise ValueError("lower_bound is required when safe_gate=True")
g = lower_bound * torch.sigmoid(rate * gate_input)
else:
g = -rate * F.softplus(gate_input)
return naive_chunk_kda(
q,
k,
v,
g,
beta,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
chunk_size=_reference_chunk_size(q.shape[1], chunk_size),
)
+5
View File
@@ -0,0 +1,5 @@
"""Incremental recurrent KDA implementations and state containers."""
from .fused import KDAState, fused_recurrent_kda, fused_recurrent_kda_step
__all__ = ["KDAState", "fused_recurrent_kda", "fused_recurrent_kda_step"]
+73
View File
@@ -0,0 +1,73 @@
"""L6: FLA fused recurrent KDA decode with optional step cache."""
from __future__ import annotations
from dataclasses import dataclass
import torch
from kda._fla.ops.kda.fused_recurrent import fused_recurrent_kda as _fused_recurrent_kda
@dataclass
class KDAState:
"""Mutable recurrent state cache: ``S`` is ``[B, HV, K, V]``."""
S: torch.Tensor
pos: int = 0
def reset(self):
self.S.zero_()
self.pos = 0
def fused_recurrent_kda_step(
state: KDAState,
q_t: torch.Tensor,
k_t: torch.Tensor,
v_t: torch.Tensor,
g_t: torch.Tensor,
beta_t: torch.Tensor,
scale: float | None = None,
):
"""Single-token step. Inputs are ``[B, H|HV, ...]`` (no time dim)."""
o, ht = _fused_recurrent_kda(
q_t.unsqueeze(1),
k_t.unsqueeze(1),
v_t.unsqueeze(1),
g_t.unsqueeze(1),
beta_t.unsqueeze(1),
scale=scale,
initial_state=state.S,
output_final_state=True,
)
state.S = ht
state.pos += 1
return o.squeeze(1)
def fused_recurrent_kda(
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,
**kwargs,
):
return _fused_recurrent_kda(
q,
k,
v,
g,
beta,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
**kwargs,
)
__all__ = ["KDAState", "fused_recurrent_kda", "fused_recurrent_kda_step"]
+13
View File
@@ -0,0 +1,13 @@
"""Readable PyTorch implementations used as correctness references."""
from .chunkwise import naive_chunk_kda
from .gate import kda_gate_naive, kda_gate_reference
from .recurrent import naive_kda, naive_kda_fwd
__all__ = [
"kda_gate_naive",
"kda_gate_reference",
"naive_chunk_kda",
"naive_kda",
"naive_kda_fwd",
]
+155
View File
@@ -0,0 +1,155 @@
"""Pure-PyTorch chunked reference implementation of KDA."""
from __future__ import annotations
import math
import warnings
import torch
from einops import rearrange
#: ``exp`` overflows past this exponent in fp32 and bf16 (both top out at 3.4e38).
_EXP_LIMIT = math.log(torch.finfo(torch.float32).max)
#: Row-block size for :func:`_decayed_dot`.
#:
#: The g_ref GEMM exponentiates the gate span between the reference row and the
#: rows/columns it covers, so the block size caps that exponent at
#: ``DECAY_BLOCK * max|g|``. With the default ``lower_bound=-5`` gate that is
#: ``16 * 5 = 80 < ln(3.4e38) = 88.7``, i.e. fp32/bf16-safe for any chunk size.
#: Referencing a whole 64-row chunk instead would allow ``64 * 5 = 320`` and
#: overflow to NaN once the gate saturates.
DECAY_BLOCK = 16
def _decayed_dot(x: torch.Tensor, k: torch.Tensor, g: torch.Tensor) -> torch.Tensor:
"""Return ``A[..., i, j] = <x_i, exp(g_i-g_j) * k_j>`` (FLA g_ref GEMM).
Only the causal part (``j <= i``) is exact; callers mask the rest, which is
left at zero. Rows are processed in blocks of :data:`DECAY_BLOCK` against
the block's own first row, which is what bounds the exponent: for a row
block starting at ``r``, ``exp(g_i - g_ref)`` spans at most ``DECAY_BLOCK``
steps, and ``exp(g_ref - g_j)`` is ``<= 1`` for ``j < r`` and likewise spans
at most ``DECAY_BLOCK`` steps for ``j >= r``.
"""
C = g.shape[-2]
out = g.new_zeros(*g.shape[:-1], C)
for r in range(0, C, DECAY_BLOCK):
end = min(r + DECAY_BLOCK, C)
g_ref = g[..., r : r + 1, :]
rows = x[..., r:end, :] * (g[..., r:end, :] - g_ref).exp()
cols = k[..., :end, :] * (g_ref - g[..., :end, :]).exp()
out[..., r:end, :end] = rows @ cols.transpose(-1, -2)
return out
#: Whether :func:`naive_chunk_kda` checks the gate span against the ``exp``
#: budget. The check costs one device sync per call; set it to ``False`` if that
#: matters more than diagnosing a NaN.
CHECK_DECAY_SPAN = True
def _max_decay_span(g_cumsum: torch.Tensor) -> torch.Tensor:
"""Largest ``|g_ref - g_j|`` any row block will exponentiate."""
C = g_cumsum.shape[-2]
if C % DECAY_BLOCK == 0:
blocks = g_cumsum.unflatten(-2, (C // DECAY_BLOCK, DECAY_BLOCK))
return (blocks[..., :1, :] - blocks).abs().amax()
return torch.stack(
[
(g_cumsum[..., r : r + 1, :] - g_cumsum[..., r : r + DECAY_BLOCK, :])
.abs()
.amax()
for r in range(0, C, DECAY_BLOCK)
]
).amax()
def _warn_if_decay_span_overflows(g_cumsum: torch.Tensor) -> None:
"""Warn when a row block's gate span is about to overflow ``exp``.
``DECAY_BLOCK`` bounds this for the default ``safe_gate`` path, but an
unbounded gate (``-A.exp() * softplus(x)``) can still exceed it.
"""
span = _max_decay_span(g_cumsum).item()
if span > _EXP_LIMIT:
warnings.warn(
f"gate span within a {DECAY_BLOCK}-row block is {span:.1f} > "
f"{_EXP_LIMIT:.1f}; exp() will overflow to inf and the output will "
"be NaN. Reduce the gate magnitude (e.g. safe_gate with a smaller "
"|lower_bound|) or use backend='triton'.",
RuntimeWarning,
stacklevel=3,
)
def naive_chunk_kda(
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,
):
"""Chunk-parallel, inter-chunk recurrent KDA reference.
Shapes are ``q/k: [B,T,H,K]``, ``v: [B,T,HV,V]``,
``g: [B,T,HV,K]`` and ``beta: [B,T,HV]``.
"""
dtype = v.dtype
B, T, H, K = q.shape
HV, V = v.shape[2], v.shape[-1]
C = chunk_size
assert HV % H == 0, f"HV={HV} must be divisible by H={H}"
assert T % C == 0, f"T={T} must be divisible by chunk_size={C}"
scale = K**-0.5 if scale is None else scale
q, k = [
rearrange(x, "b (n c) h d -> b h n c d", c=C)
.repeat_interleave(HV // H, dim=1)
for x in (q, k)
]
v, g = [rearrange(x, "b (n c) h d -> b h n c d", c=C) for x in (v, g)]
beta = rearrange(beta, "b (n c) h -> b h n c", c=C)
q = q * scale
g = g.cumsum(dim=-2)
if CHECK_DECAY_SPAN:
_warn_if_decay_span_overflows(g)
# r_i + sum_{j<i} beta_j <k_i, exp(g_i-g_j)k_j> r_j
# = v_i - <exp(g_i)k_i, S_start>.
mask_upper = torch.triu(torch.ones(C, C, dtype=torch.bool, device=q.device))
mask_strict_upper = torch.triu(mask_upper, diagonal=1)
eye = torch.eye(C, dtype=q.dtype, device=q.device)
A_kk = _decayed_dot(k, k, g)
M = eye + (A_kk * beta[..., None, :]).masked_fill(mask_upper, 0)
W = torch.linalg.solve_triangular(M, g.exp() * k, upper=False)
U = torch.linalg.solve_triangular(M, v, upper=False)
# Output includes the current token, hence the diagonal is retained.
A_qk = (_decayed_dot(q, k, g) * beta[..., None, :]).masked_fill(mask_strict_upper, 0)
S = q.new_zeros(B, HV, K, V)
if initial_state is not None:
S = S + initial_state
o = v.new_empty(B, HV, T // C, C, V)
for n in range(T // C):
q_n, k_n, g_n = q[:, :, n], k[:, :, n], g[:, :, n]
r = U[:, :, n] - W[:, :, n] @ S
o[:, :, n] = (q_n * g_n.exp()) @ S + A_qk[:, :, n] @ r
decay = (g_n[:, :, -1:, :] - g_n).exp()
S = S * g_n[:, :, -1, :, None].exp()
S = S + (decay * k_n).transpose(-1, -2) @ (r * beta[:, :, n, :, None])
if not output_final_state:
S = None
return rearrange(o, "b h n c d -> b (n c) h d").to(dtype), S
# Backward-compatible name used by earlier notes/scripts.
naive_chunk_kda_fwd = naive_chunk_kda
+47
View File
@@ -0,0 +1,47 @@
"""PyTorch references for the two KDA gate activations."""
from __future__ import annotations
import torch
import torch.nn.functional as F
def kda_gate_reference(
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None = None,
*,
safe_gate: bool = False,
lower_bound: float | None = None,
) -> torch.Tensor:
"""Compute the official KDA gate semantics in PyTorch.
``A_log`` is head-wise with shape ``[HV]`` and ``dt_bias`` is
per-dimension with shape ``[HV, K]`` (or flattened to ``[HV*K]``).
"""
HV, K = g.shape[-2:]
gate_input = g if dt_bias is None else g + dt_bias.view(HV, K)
rate = A_log.view(HV, 1).exp()
if safe_gate:
if lower_bound is None:
raise ValueError("lower_bound is required when safe_gate=True")
return lower_bound * torch.sigmoid(rate * gate_input)
return -rate * F.softplus(gate_input)
def kda_gate_naive(
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None = None,
lower_bound: float | None = None,
) -> torch.Tensor:
"""Compatibility name matching FLA's reference gate convention."""
return kda_gate_reference(
g,
A_log,
dt_bias,
safe_gate=lower_bound is not None,
lower_bound=lower_bound,
)
__all__ = ["kda_gate_naive", "kda_gate_reference"]
+298
View File
@@ -0,0 +1,298 @@
"""L1: Naive recurrent KDA fwd+bwd (torch only).
公式 (per timestep t, log-space gate; q/k 入口 H 维, 内部 repeat_interleave 到 HV):
S_t = exp(g_t) * S_{t-1} + (beta_t * k_t) outer (v_t - k_t . (exp(g_t) * S_{t-1}))
o_t = (q_t * scale) . S_t
backward (BPTT, T -> 0):
设 dS_t 为进入 t 步累积的反传梯度 (含 o_t 反传).
1. o_t = q_t . S_t -> dS_t += q_t outer do_t (i.e. dS = dS + q_t·do_t)
dq_t = do_t . S_t^T -> einsum('bhv,bhkv->bhk')
2. S_t = S_decay + a_t outer r_t, a_t = b_t k_t, r_t = v_t - k_t . S_decay
其中 S_decay = exp(g_t) * S_{t-1}
dS_{t-1} = exp(g_t) * (dS_t - r_t outer da_t - a_t outer dr_t) via residual 反传
更具体:
dS_decay = dS_t - (a_t outer dr_t) - (da_t outer r_t)
dS_{t-1} += exp(g_t) * dS_decay
其中 dr_t = -dv_t + dS_t . a_t^T (因为 r_t = v - k·S_dec, dr 来自 -dv - k·dS_decay)
da_t = -r_t outer dS_t? 让我直接推导下面.
推导 (设 G1 = S_t, 走 a = r 反向链 通过 autograd):
o_t = q_t . G1
dq_t = do_t . G1^T -> [B,HV,K]
dG1 = q_t outer do_t -> [B,HV,K,V] = dS_t (上游)
G1 = Sdec + a outer r -> Sdec = G1[...] (跳过)
dSdec = dG1
da_t = r_t outer dG1 -> [B,HV,K] (因为 a outer r 是 K-V, d(a outer r) = r outer d[...,V])
但在 einsum 表示: dA_t.grad = einsum('bhkv,bhv->bhk', dS_t, r_t)
dr_t = a_t outer dG1 -> [B,HV,V] = einsum('bhkv,bhk->bhv', dS_t, a_t)
plus: S_t = Sdec + a outer r -> a outer r - outer product 形状是 [B,HV,K,V] = einsum('bhk,bhv->bhkv')
d(a outer r) 的雅可比: let G1_m = a_t ⊗ r_t (rank-1 matrix per (b,h))
dG1_m[i,j] = da_t[i] * r_t[j] + a_t[i] * dr_t[j]
在外积形式, 即 dG1_m = a_outer r 的张量积正交分解:
da_t = sum_j r_t[j] dG1_m[i,j] = einsum('bhkv,bhv->bhk', dG1_m, r_t)
dr_t = sum_i a_t[i] dG1_m[i,j] = einsum('bhkv,bhk->bhv', dG1_m, a_t)
因为 a_t = b_t k_t -> da_t = db_t k_t + b_t dk_t (b_t 是 ...)
db_t = einsum('bhk,bhk->bh', da_t, k_t)
dk_t_a = b_t * da_t (来自 a_t 路径, 还有来自 r_t 路径和 S_dec 路径)
因为 r_t = v_t - k_t . S_dec -> 注 rk_t grad via dg,S_dec 和 dv_t
dv_t = -dr_t (实际 dr 的负梯度) 即 dv_t = -dr_t
这里 r_t = v_t - k_t · S_dec, 写作矩阵乘 r = v - einsum('bhk,bhkv->bhv', k, S_dec)
dr = -dv - einsum('bhk,bhkv->bhv', dk_from_r, S_dec) + eigengrad via S_dec
更精确的反向: r_t = v_t - k_t . S_dec
dv_t += -dr_t -> dv_t = -dr_t
dk_t_r_path = -S_dec outer dr_t (即 -dS_dec 传递来自 k_t 的部分)
具体: d(k·S) = dk·S + k·dS -> dS_dec 这层, dk 的贡献: -S_dec outer dr_t
即 dk_t_r = einsum('bhv,bhkv->bhk', -dr_t, S_dec)
dS_dec_r = -k_t outer dr_t = -einsum('bhv,bhk->bhkv', dr_t, k_t)
合并: dS_dec 合总 = dG1 + (-k_t outer dr_t)
= dS_t - k_t outer dr_t
(相加过的 dv, dk_r, dS_dec_r 都上面项)
Sdec = exp(g_t) * S_{t-1}:
dS_{t-1} = exp(g_t) ⊙ dS_dec (因为 Sdec = exp_g * S_prev, 微分后 exp_g 直接相乘)
dg_t = exp(g_t) * S_prev * dS_dec (微分时对 g_t (log-space) 求偏导数)
即 dg_t = exp(g_t) * (S_{t-1} ⊙ dS_dec) -> 沿 K 维求和
in einsum: dg_t = sum over v of (exp(g_t) * S_{t-1}) ⊙ dS_dec ...\n
= einsum('bhk, bhk, bhkv -> bhk', exp_g, S_prev, dS_dec)
更简洁: Sdec = exp_g * S_prev (per-(b,h,k)/v), 故 dSdec/dg_t = S_prev * exp_g
所以 dg_t = sum_v S_prev_sub_k_dim * exp_g * dS_dec -> [B, HV, K]
einsum: dg_t = einsum('bhkv,bhkv->bhk', Sdec, dS_dec)
(因为 Sdec = S_prev * exp_g, sum_v Sdec[:, :, :, v] * dS_dec[:, :, :, v] = sum_v Sdec_eachK * dSdec_eachK)
einsum上是 einsum('bhkv,bhkv->bhk', Sdec, dSdec)
dS_{t-1} = exp_g ⊙ dSdec (per (b,h,k,v) entrywise multiply exp_g with dSdec)
GVA 反归约:
q,k 入口 [B, T, H, K] --repeat_interleave(G, dim=2)--> [B, T, HV, K]
内部计算后, dq/dk 在 HV 维上 -> dV 拿 shape [B,T,HV,K]
bwd 通过 sum 回 H: dq_H = dq_HV.view(B,T,H,G,K).sum(dim=3) -> [B,T,H,K]
(因为 repeat_interleave 是复制, 反传是 sum 路径相同意义)
记号对照:
a_t = b_t * k_t (a = beta * k) [B, HV, K]
r_t = v_t - k_t . S_dec (residual) [B, HV, V]
S_dec = exp(g_t) * S_{t-1} [B, HV, K, V]
S_t = S_dec + a_t outer r_t [B, HV, K, V]
o_t = q_t . S_t = (q_t_eff * scale) . S_t [B, HV, V]
"""
from __future__ import annotations
import math
import torch
def naive_kda_fwd(
q: torch.Tensor, # [B, T, H, K]
k: torch.Tensor, # [B, T, H, K]
v: torch.Tensor, # [B, T, HV, V]
g: torch.Tensor, # [B, T, HV, K]
beta: torch.Tensor, # [B, T, HV]
scale: float | None = None,
initial_state: torch.Tensor | None = None, # [B, HV, K, V]
output_final_state: bool = False,
*,
force_float32: bool = False,
):
"""纯 forward, 不带 autograd. 与上游 naive_recurrent_kda 数值等价.
force_float32=True 时强制 fp32 计算 (与上游对拍时用);
默认保持输入 dtype (gradcheck 用 fp64).
"""
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 = 1.0 / math.sqrt(K)
# 上游强制 fp32; 本实现默认保留输入 dtype 以便 gradcheck 适用 fp64
# force_float32=True 时与上游逐位对齐
work_dtype = torch.float if force_float32 else q.dtype
q = q.to(work_dtype)
k = k.to(work_dtype)
v = v.to(work_dtype)
g = g.to(work_dtype)
beta = beta.to(work_dtype)
# GVA: expand q/k from H to HV
qe = q.repeat_interleave(G, dim=2) * scale # [B, T, HV, K]
ke = k.repeat_interleave(G, dim=2) # [B, T, HV, K]
S = torch.zeros(B, HV, K, V, dtype=work_dtype, device=q.device)
if initial_state is not None:
S = S + initial_state.to(work_dtype)
o = torch.empty(B, T, HV, V, dtype=work_dtype, device=q.device)
for t in range(T):
q_t = qe[:, t] # [B, HV, K]
k_t = ke[:, t] # [B, HV, K]
v_t = v[:, t] # [B, HV, V]
g_t = g[:, t] # [B, HV, K]
b_t = beta[:, t] # [B, HV]
S_dec = S * g_t.exp().unsqueeze(-1) # [B, HV, K, V]
p_t = torch.einsum('b h k, b h k v -> b h v', k_t, S_dec) # [B, HV, V]
r_t = v_t - p_t # [B, HV, V]
a_t = b_t.unsqueeze(-1) * k_t # [B, HV, K]
S = S_dec + torch.einsum('b h k, b h v -> b h k v', a_t, r_t)
o[:, t] = torch.einsum('b h k, b h k v -> b h v', q_t, S)
if not output_final_state:
S = None
return o.to(dtype), S
class KDAFunction(torch.autograd.Function):
"""autograd Function (forward + backward).
forward 入参顺序 (q, k, v, g, beta, scale, initial_state, output_final_state)
backward 必须返回一致: (dq, dk, dv, dg, dbeta, None, dinit_state, None)
"""
@staticmethod
def forward(ctx, q, k, v, g, beta, scale, initial_state, output_final_state):
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 = 1.0 / math.sqrt(K)
work_dtype = q.dtype
qf = q.to(work_dtype).contiguous()
kf = k.to(work_dtype).contiguous()
vf = v.to(work_dtype).contiguous()
gf = g.to(work_dtype).contiguous()
bf = beta.to(work_dtype).contiguous()
# GVA: expand q/k from H to HV
qe = qf.repeat_interleave(G, dim=2) * scale # [B, T, HV, K]
ke = kf.repeat_interleave(G, dim=2) # [B, T, HV, K]
S = torch.zeros(B, HV, K, V, dtype=work_dtype, device=q.device)
if initial_state is not None:
S = S + initial_state.to(work_dtype)
o = torch.empty(B, T, HV, V, dtype=work_dtype, device=q.device)
q_ts, k_ts, b_ts, S_decs, r_ts, a_ts, exp_g_ts = [], [], [], [], [], [], []
for t in range(T):
q_t = qe[:, t]
k_t = ke[:, t]
v_t = vf[:, t]
g_t = gf[:, t]
b_t = bf[:, t]
exp_g_t = g_t.exp()
S_dec = S * exp_g_t.unsqueeze(-1)
p_t = torch.einsum('b h k, b h k v -> b h v', k_t, S_dec)
r_t = v_t - p_t
a_t = b_t.unsqueeze(-1) * k_t
S = S_dec + torch.einsum('b h k, b h v -> b h k v', a_t, r_t)
o[:, t] = torch.einsum('b h k, b h k v -> b h v', q_t, S)
q_ts.append(q_t)
k_ts.append(k_t)
b_ts.append(b_t)
S_decs.append(S_dec)
r_ts.append(r_t)
a_ts.append(a_t)
exp_g_ts.append(exp_g_t)
ctx.save_for_backward(
torch.stack(q_ts, dim=1),
torch.stack(k_ts, dim=1),
torch.stack(b_ts, dim=1),
torch.stack(S_decs, dim=1),
torch.stack(r_ts, dim=1),
torch.stack(a_ts, dim=1),
torch.stack(exp_g_ts, dim=1),
)
ctx.G = G
ctx.H = H
ctx.HV = HV
ctx.K = K
ctx.V = V
ctx.T = T
ctx.B = B
ctx.scale = scale
ctx.dtype = dtype
ctx.has_initial_state = initial_state is not None
ctx.output_final_state = output_final_state
final_S = S if output_final_state else None
return o.to(dtype), final_S
@staticmethod
def backward(ctx, do, dS):
q_ts, k_ts, b_ts, S_decs, r_ts, a_ts, exp_g_ts = ctx.saved_tensors
B, T, H, HV, K, V, G = ctx.B, ctx.T, ctx.H, ctx.HV, ctx.K, ctx.V, ctx.G
work_dtype = q_ts.dtype
device = q_ts.device
dq_e = torch.zeros(B, T, HV, K, dtype=work_dtype, device=device)
dk_e = torch.zeros(B, T, HV, K, dtype=work_dtype, device=device)
dv = torch.zeros(B, T, HV, V, dtype=work_dtype, device=device)
dg = torch.zeros(B, T, HV, K, dtype=work_dtype, device=device)
dbeta= torch.zeros(B, T, HV, dtype=work_dtype, device=device)
if dS is None:
dS_acc = torch.zeros(B, HV, K, V, dtype=work_dtype, device=device)
else:
dS_acc = dS.to(work_dtype).clone()
for t in range(T - 1, -1, -1):
q_t = q_ts[:, t]
k_t = k_ts[:, t]
b_t = b_ts[:, t]
S_dec = S_decs[:, t]
r_t = r_ts[:, t]
a_t = a_ts[:, t]
exp_g_t = exp_g_ts[:, t]
do_t = do[:, t].to(work_dtype)
S_t = S_dec + torch.einsum('b h k, b h v -> b h k v', a_t, r_t)
dS_acc = dS_acc + torch.einsum('b h k, b h v -> b h k v', q_t, do_t)
dq_e[:, t] = torch.einsum('b h v, b h k v -> b h k', do_t, S_t)
da_t = torch.einsum('b h v, b h k v -> b h k', r_t, dS_acc)
dr_t = torch.einsum('b h k, b h k v -> b h v', a_t, dS_acc)
dbeta[:, t] = torch.einsum('b h k, b h k -> b h', k_t, da_t)
dk_t_a = b_t.unsqueeze(-1) * da_t
dv[:, t] = dr_t
dS_dec_from_r = -torch.einsum('b h v, b h k -> b h k v', dr_t, k_t)
dk_t_r = -torch.einsum('b h v, b h k v -> b h k', dr_t, S_dec)
dS_dec_total = dS_acc + dS_dec_from_r
dk_e[:, t] = dk_t_a + dk_t_r
dg[:, t] = torch.einsum('b h k v, b h k v -> b h k', S_dec, dS_dec_total)
dS_acc = exp_g_t.unsqueeze(-1) * dS_dec_total
if HV > H:
dq_H = dq_e.view(B, T, H, G, K).sum(dim=3)
dk_H = dk_e.view(B, T, H, G, K).sum(dim=3)
else:
dq_H = dq_e
dk_H = dk_e
# q 在 forward 内被乘过 scale (qe = q * scale), chain rule: dq_orig = dq_e * scale
dq_H = dq_H * ctx.scale
return (dq_H.to(ctx.dtype), dk_H.to(ctx.dtype), dv.to(ctx.dtype),
dg.to(ctx.dtype), dbeta.to(ctx.dtype), None, None, None)
def naive_kda(
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,
):
"""对外入口: 调 KDAFunction.apply."""
return KDAFunction.apply(q, k, v, g, beta, scale, initial_state, output_final_state)
+7
View File
@@ -0,0 +1,7 @@
"""Local Triton KDA kernels vendored from FLA chunk_{fwd,intra,bwd,wy,gate}."""
from .chunk import ChunkKDAFunction, chunk_kda
from .chunk_fwd import chunk_kda_fwd
from .gate import kda_gate_fwd
__all__ = ["ChunkKDAFunction", "chunk_kda", "chunk_kda_fwd", "kda_gate_fwd"]
+5
View File
@@ -0,0 +1,5 @@
"""FLA ``chunk_kda`` surface used by ``ops.api`` backend='triton'."""
from kda._fla.ops.kda.chunk import ChunkKDAFunction, chunk_kda
__all__ = ["ChunkKDAFunction", "chunk_kda"]
+5
View File
@@ -0,0 +1,5 @@
"""Vendored FLA chunk KDA backward."""
from kda._fla.ops.kda.chunk_bwd import chunk_kda_bwd
__all__ = ["chunk_kda_bwd"]
+37
View File
@@ -0,0 +1,37 @@
"""Vendored FLA chunk KDA forward, returning ``(o, ht)`` like the public op."""
from __future__ import annotations
import torch
from kda._fla.ops.kda.chunk import chunk_kda
from kda._fla.ops.kda.chunk_fwd import chunk_kda_fwd as fla_chunk_kda_fwd
__all__ = ["chunk_kda_fwd", "fla_chunk_kda_fwd"]
def chunk_kda_fwd(
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,
**kwargs,
):
"""Chunked KDA forward with FLA kernels. Returns ``(o, ht)``."""
return chunk_kda(
q,
k,
v,
g,
beta,
scale=scale,
initial_state=initial_state,
output_final_state=output_final_state,
chunk_size=chunk_size,
**kwargs,
)
+36
View File
@@ -0,0 +1,36 @@
"""Vendored FLA KDA gate fusion (standard + safe gate + chunk cumsum)."""
from __future__ import annotations
import torch
from kda._fla.ops.kda.gate import (
kda_gate_bwd,
kda_gate_chunk_cumsum,
kda_gate_fwd as _kda_gate_fwd,
)
DEFAULT_LOWER_BOUND = -5.0
def kda_gate_fwd(
g: torch.Tensor,
A_log: torch.Tensor,
dt_bias: torch.Tensor | None = None,
lower_bound: float | None = DEFAULT_LOWER_BOUND,
):
return _kda_gate_fwd(
g,
A_log=A_log,
dt_bias=dt_bias,
lower_bound=lower_bound,
output_dtype=g.dtype,
)
__all__ = [
"DEFAULT_LOWER_BOUND",
"kda_gate_bwd",
"kda_gate_chunk_cumsum",
"kda_gate_fwd",
]
+5
View File
@@ -0,0 +1,5 @@
"""Vendored FLA WY recompute used by the chunk KDA backward."""
from kda._fla.ops.kda.wy_fast import recompute_w_u_fwd
__all__ = ["recompute_w_u_fwd"]