"""L6: fused_recurrent decode matches naive recurrent.""" import pytest import torch from kda.ops.recurrent.fused import fused_recurrent_kda from kda.ops.reference.recurrent import naive_kda pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") def test_step_matches_naive(): B, T, H, HV, K, V = 2, 32, 2, 4, 8, 8 torch.manual_seed(20) device = "cuda" q = torch.randn(B, T, H, K, device=device) k = torch.randn(B, T, H, K, device=device) v = torch.randn(B, T, HV, V, device=device) g = -torch.rand(B, T, HV, K, device=device) * 2 beta = torch.rand(B, T, HV, device=device) o_naive, _ = naive_kda( q.double(), k.double(), v.double(), g.double(), beta.double(), output_final_state=False, ) o_step, _ = fused_recurrent_kda(q, k, v, g, beta, output_final_state=False) torch.testing.assert_close(o_step.float(), o_naive.float(), rtol=2e-3, atol=2e-3) if __name__ == "__main__": test_step_matches_naive()