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

73 lines
2.8 KiB
Python

"""严重度扫描: 把"数学爆炸"和"两条路径分道扬镳"分开测。
对每个 k 缩放系数 s:
rec64 = 逐步递推 float64 (无三角求解) -> 数学真值
chk64 = chunkwise float64 (全局 solve)
chk32 = chunkwise float32
tri32 = FLA/vendored triton 16x16 分块路径 (float32 in/out)
"""
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
dev = "cuda"
B, T, H, K, V = 1, 64, 1, 16, 32
try:
from kda.ops.triton.chunk import chunk_kda as triton_chunk_kda
except Exception as e: # pragma: no cover
triton_chunk_kda = None
print("triton backend unavailable:", e)
def make(scale_k, dtype):
gen = torch.Generator(device=dev).manual_seed(0)
q = torch.randn(B, T, H, K, device=dev, dtype=dtype, generator=gen)
k = torch.randn(B, T, H, K, device=dev, dtype=dtype, generator=gen)
v = torch.randn(B, T, H, V, device=dev, dtype=dtype, generator=gen)
g = -F.softplus(torch.randn(B, T, H, K, device=dev, dtype=dtype, generator=gen)) * 0.1
beta = torch.rand(B, T, H, device=dev, dtype=dtype, generator=gen)
return q, F.normalize(k, dim=-1) * scale_k, v, g, beta
def amax(x):
return x.abs().max().item()
hdr = (f"{'‖k‖':>6} | {'max|N|':>9} | {'‖M⁻¹‖∞':>10} | {'rec fp64':>10} | "
f"{'chk fp64':>10} | {'chk fp32':>10} | {'tri fp32':>10} | "
f"{'rel(chk64/rec64)':>16} | {'rel(chk32/chk64)':>16} | {'rel(tri32/chk32)':>16}")
print(hdr)
print("-" * len(hdr))
for s in (1.0, 1.5, 2.0, 2.5, 3.0, 4.0, 5.0):
q, k, v, g, beta = make(s, torch.float64)
o_rec, _ = naive_kda_fwd(q, k, v, g, beta)
o_c64, _ = naive_chunk_kda(q, k, v, g, beta, chunk_size=64)
q3, k3, v3, g3, b3 = [x.float() for x in (q, k, v, g, beta)]
o_c32, _ = naive_chunk_kda(q3, k3, v3, g3, b3, chunk_size=64)
if triton_chunk_kda is not None:
o_t32, _ = triton_chunk_kda(q3, k3, v3, g3, b3, chunk_size=64)
o_t32 = o_t32.double()
else:
o_t32 = torch.full_like(o_c64, float("nan"))
gc = g.cumsum(dim=1)[0, :, 0]
kk = k[0, :, 0]
gref = gc[:1]
A = (kk * (gc - gref).exp()) @ (kk * (gref - gc).exp()).T
M = torch.eye(T, dtype=torch.float64, device=dev) + (A * beta[0, :, 0][None, :]).tril(-1)
minv = torch.linalg.inv(M).abs().sum(1).max().item()
def rel(a, b):
return ((a - b).abs().max() / (b.abs().max() + 1e-300)).item()
print(f"{k.norm(dim=-1).mean():6.2f} | {amax((A * beta[0, :, 0][None, :]).tril(-1)):9.2e} | "
f"{minv:10.2e} | {amax(o_rec):10.2e} | {amax(o_c64):10.2e} | "
f"{amax(o_c32):10.2e} | {amax(o_t32):10.2e} | "
f"{rel(o_c64, o_rec):16.2e} | {rel(o_c32.double(), o_c64):16.2e} | "
f"{rel(o_t32, o_c32.double()):16.2e}")