Files
K3/kda/ops/reference/recurrent.py
T
dela 584f7e9e73 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.
2026-08-25 14:43:17 +08:00

299 lines
13 KiB
Python

"""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)