Files
K3/notes/verify_l2norm_stability.py
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

84 lines
3.5 KiB
Python

"""验证: 不做 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))