Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
73 lines
2.8 KiB
Python
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}")
|