"""L1: gradcheck for naive_recurrent_kda. 验证策略: torch.autograd.gradcheck 走 forward+backward 五个梯度. 强制 dtype=float64; eps=1e-6, atol=1e-4. shape (small): B=2, T=8, H=2, HV=4, K=4, V=4 """ import torch from kda.ops.reference.recurrent import naive_kda def test_gradcheck(): B, T, H, HV, K, V = 2, 8, 2, 4, 4, 4 torch.manual_seed(0) # 所有输入都需要 requires_grad=True 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_kda(q, k, v, g, b, output_final_state=True), (q, k, v, g, beta), eps=1e-6, atol=1e-4, rtol=1e-3, ), "L1 gradcheck 失败" print("L1 gradcheck: PASSED") if __name__ == "__main__": test_gradcheck()