"""L4: vendored FLA Triton bwd vs naive chunked. Triton kernels run in fp32, so this checks VJP vs L2 rather than fp64 gradcheck. Inputs L2-normalize q/k like the trained KDA path. """ import pytest import torch import torch.nn.functional as F from kda.ops.reference.chunkwise import naive_chunk_kda from kda.ops.triton.chunk_fwd import chunk_kda_fwd pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") def test_bwd_matches_naive(): B, T, H, HV, K, V = 2, 64, 4, 8, 32, 32 torch.manual_seed(5) dev = "cuda" q = F.normalize(torch.randn(B, T, H, K, device=dev), dim=-1).requires_grad_() k = F.normalize(torch.randn(B, T, H, K, device=dev), dim=-1).requires_grad_() v = torch.randn(B, T, HV, V, device=dev, requires_grad=True) g = (-torch.rand(B, T, HV, K, device=dev) * 2).requires_grad_() beta = torch.rand(B, T, HV, device=dev, requires_grad=True) o_na, _ = naive_chunk_kda(q, k, v, g, beta, chunk_size=64) o_na.sum().backward() dq_na, dk_na = q.grad.clone(), k.grad.clone() dv_na, dg_na, db_na = v.grad.clone(), g.grad.clone(), beta.grad.clone() q.grad = k.grad = v.grad = g.grad = beta.grad = None o_tr, _ = chunk_kda_fwd(q, k, v, g, beta, chunk_size=64) o_tr.sum().backward() dq_tr, dk_tr = q.grad.clone(), k.grad.clone() dv_tr, dg_tr, db_tr = v.grad.clone(), g.grad.clone(), beta.grad.clone() torch.testing.assert_close(dq_tr, dq_na, rtol=2e-2, atol=2e-3) torch.testing.assert_close(dk_tr, dk_na, rtol=2e-2, atol=2e-3) torch.testing.assert_close(dv_tr, dv_na, rtol=2e-2, atol=2e-3) torch.testing.assert_close(dg_tr, dg_na, rtol=2e-2, atol=2e-3) torch.testing.assert_close(db_tr, db_na, rtol=2e-2, atol=2e-3) def test_triton_kda_attention_backward_dt_bias_rank2(): """Layer stores dt_bias as [HV, K]; Triton bwd used to return a flat [HV*K].""" from kda.layers.kda_attn import KDAAttention torch.manual_seed(0) model = KDAAttention( hidden_size=64, num_heads=4, num_value_heads=8, head_dim=16, chunk_size=64, kda_backend="triton", ).cuda() x = torch.randn(2, 64, 64, device="cuda") model(x).sum().backward() assert model.dt_bias.grad is not None assert model.dt_bias.grad.shape == model.dt_bias.shape assert model.A_log.grad is not None assert torch.isfinite(model.dt_bias.grad).all() if __name__ == "__main__": test_bwd_matches_naive() test_triton_kda_attention_backward_dt_bias_rank2()