"""The vendored FLA fused gate must match the PyTorch reference.""" import pytest import torch from kda.ops.reference.gate import kda_gate_reference from kda.ops.triton.gate import kda_gate_fwd pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") def test_triton_gate_matches_reference(): B, T, HV, K = 1, 32, 4, 8 torch.manual_seed(10) device = "cuda" g = torch.randn(B, T, HV, K, device=device, dtype=torch.float32) A_log = torch.randn(HV, device=device, dtype=torch.float32) * 0.5 dt_bias = torch.randn(HV, K, device=device, dtype=torch.float32) * 0.1 expected = kda_gate_reference(g, A_log, dt_bias) actual = kda_gate_fwd(g, A_log, dt_bias, lower_bound=None) torch.testing.assert_close(actual, expected, rtol=1e-4, atol=1e-4)