"""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()