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:
@@ -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)
|
||||
Reference in New Issue
Block a user