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