"""验证: 不做 L2-norm 时 KDA 的爆炸是"数学上真实"还是"浮点误差"。 判据: 逐步递推 (naive_kda_fwd, 无三角求解) 在 float64 下的输出。 - 若 fp64 递推也 ~1e32 => 爆炸是 KDA 递推本身的数学性质 - 若 fp64 递推 O(1) 而 chunkwise 爆炸 => 是 solve 的浮点失效 """ import torch import torch.nn.functional as F from kda.ops.reference.chunkwise import naive_chunk_kda from kda.ops.reference.recurrent import naive_kda_fwd torch.manual_seed(0) dev = "cuda" B, T, H, K, V = 1, 64, 1, 16, 32 C = 64 def make(norm: bool, dtype): g0 = torch.Generator(device=dev).manual_seed(0) q = torch.randn(B, T, H, K, device=dev, dtype=dtype, generator=g0) k = torch.randn(B, T, H, K, device=dev, dtype=dtype, generator=g0) v = torch.randn(B, T, H, V, device=dev, dtype=dtype, generator=g0) # 训练里 g = -A.exp()*softplus(...) < 0,量级温和 g = -F.softplus(torch.randn(B, T, H, K, device=dev, dtype=dtype, generator=g0)) * 0.1 beta = torch.rand(B, T, H, device=dev, dtype=dtype, generator=g0) if norm: q, k = F.normalize(q, dim=-1), F.normalize(k, dim=-1) return q, k, v, g, beta def amax(x): return x.abs().max().item() print(f"{'norm':>5} | {'‖k‖':>6} | {'rec fp64':>10} | {'rec fp32':>10} | " f"{'chunk fp64':>10} | {'chunk fp32':>10} | {'rel(chunk64,rec64)':>18}") print("-" * 100) for norm in (True, False): row = {} for dt in (torch.float64, torch.float32): q, k, v, g, beta = make(norm, dt) o_rec, _ = naive_kda_fwd(q, k, v, g, beta) o_chk, _ = naive_chunk_kda(q, k, v, g, beta, chunk_size=C) row[dt] = (amax(o_rec), amax(o_chk), o_chk, o_rec) knorm = make(norm, torch.float64)[1].norm(dim=-1).mean().item() r64, c64, oc64, or64 = row[torch.float64] r32, c32, oc32, _ = row[torch.float32] rel = ((oc64 - or64).abs().max() / (or64.abs().max() + 1e-30)).item() print(f"{str(norm):>5} | {knorm:6.2f} | {r64:10.3e} | {r32:10.3e} | " f"{c64:10.3e} | {c32:10.3e} | {rel:18.3e}") # ---- M 的结构诊断 ---- print("\n=== M = I + tril(A_kk*beta, -1) 诊断 (float64) ===") print(f"{'norm':>5} | {'max|N|':>9} | {'‖M⁻¹‖∞':>10} | {'cond2(M)':>10} | {'|S_final|max':>12}") for norm in (True, False): q, k, v, g, beta = make(norm, torch.float64) gc = g.cumsum(dim=1)[0, :, 0] # [T,K] kk = k[0, :, 0] # [T,K] gref = gc[:1] A = (kk * (gc - gref).exp()) @ (kk * (gref - gc).exp()).T N = (A * beta[0, :, 0][None, :]).tril(-1) M = torch.eye(T, dtype=torch.float64, device=dev) + N Minv = torch.linalg.inv(M) _, Sf = naive_kda_fwd(q, k, v, g, beta, output_final_state=True) print(f"{str(norm):>5} | {amax(N):9.3e} | {Minv.abs().sum(1).max():10.3e} | " f"{torch.linalg.cond(M).item():10.3e} | {amax(Sf):12.3e}") # ---- ‖M⁻¹‖ 随 chunk 长度的增长 ---- print("\n=== ‖M⁻¹‖∞ vs chunk 长度 C (float64) ===") for norm in (True, False): q, k, v, g, beta = make(norm, torch.float64) gc = g.cumsum(dim=1)[0, :, 0] kk = k[0, :, 0] out = [] for C_ in (4, 8, 16, 32, 64): gs, ks, bs = gc[:C_], kk[:C_], beta[0, :C_, 0] gref = gs[:1] A = (ks * (gs - gref).exp()) @ (ks * (gref - gs).exp()).T M = torch.eye(C_, dtype=torch.float64, device=dev) + (A * bs[None, :]).tril(-1) out.append(f"C={C_:2d}:{torch.linalg.inv(M).abs().sum(1).max():.2e}") print(f" norm={str(norm):>5} " + " ".join(out))