Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
48 lines
1.8 KiB
Python
48 lines
1.8 KiB
Python
"""L2: chunked naive vs L1 naive (fwd) + gradcheck."""
|
|
import torch
|
|
|
|
from kda.ops.reference.chunkwise import naive_chunk_kda
|
|
from kda.ops.reference.recurrent import naive_kda
|
|
|
|
|
|
def test_fwd_matches_naive():
|
|
"""chunked vs naive fwd, atol=1e-4."""
|
|
B, T, H, HV, K, V = 2, 32, 2, 4, 8, 8
|
|
torch.manual_seed(1)
|
|
q = torch.randn(B, T, H, K, dtype=torch.float64)
|
|
k = torch.randn(B, T, H, K, dtype=torch.float64)
|
|
v = torch.randn(B, T, HV, V, dtype=torch.float64)
|
|
g = torch.randn(B, T, HV, K, dtype=torch.float64) * 0.1
|
|
beta = torch.rand(B, T, HV, dtype=torch.float64)
|
|
|
|
o_ref, _ = naive_kda(q, k, v, g, beta, output_final_state=True)
|
|
o_chk, _ = naive_chunk_kda(q, k, v, g, beta,
|
|
output_final_state=True, chunk_size=8)
|
|
|
|
diff = (o_ref - o_chk).abs().max().item()
|
|
assert diff < 1e-4, f"chunk vs naive fwd max diff {diff:.2e} > 1e-4"
|
|
print(f"L2 fwd-vs-naive: PASSED (max diff {diff:.2e})")
|
|
|
|
|
|
def test_gradcheck():
|
|
"""L2 chunked gradcheck atol=1e-4 (allowing rtol 1e-3 for triangular path)."""
|
|
B, T, H, HV, K, V = 2, 16, 2, 4, 4, 4
|
|
torch.manual_seed(2)
|
|
q = torch.randn(B, T, H, K, dtype=torch.float64, requires_grad=True)
|
|
k = torch.randn(B, T, H, K, dtype=torch.float64, requires_grad=True)
|
|
v = torch.randn(B, T, HV, V, dtype=torch.float64, requires_grad=True)
|
|
g = torch.randn(B, T, HV, K, dtype=torch.float64, requires_grad=True) * 0.1
|
|
beta = torch.rand(B, T, HV, dtype=torch.float64, requires_grad=True)
|
|
|
|
assert torch.autograd.gradcheck(
|
|
lambda q, k, v, g, b: naive_chunk_kda(q, k, v, g, b, chunk_size=4)[0],
|
|
(q, k, v, g, beta),
|
|
eps=1e-6, atol=1e-4, rtol=1e-3,
|
|
), "L2 gradcheck 失败"
|
|
print("L2 gradcheck: PASSED")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
test_fwd_matches_naive()
|
|
test_gradcheck()
|