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.
This commit is contained in:
@@ -0,0 +1,72 @@
|
||||
"""严重度扫描: 把"数学爆炸"和"两条路径分道扬镳"分开测。
|
||||
|
||||
对每个 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}")
|
||||
Reference in New Issue
Block a user