"""KDA gate reference semantics and autograd coverage.""" import torch import torch.nn.functional as F from kda.ops.reference.gate import kda_gate_reference def test_standard_gate_matches_formula(): B, T, HV, K = 1, 32, 4, 8 torch.manual_seed(10) g = torch.randn(B, T, HV, K, dtype=torch.float64) A_log = torch.randn(HV, dtype=torch.float64) * 0.5 dt_bias = torch.randn(HV, K, dtype=torch.float64) * 0.1 expected = -A_log[:, None].exp() * F.softplus(g + dt_bias) actual = kda_gate_reference(g, A_log, dt_bias) torch.testing.assert_close(actual, expected) def test_safe_gate_matches_formula(): B, T, HV, K = 1, 8, 2, 4 torch.manual_seed(12) g = torch.randn(B, T, HV, K, dtype=torch.float64) A_log = torch.randn(HV, dtype=torch.float64) dt_bias = torch.randn(HV, K, dtype=torch.float64) lower_bound = -5.0 expected = lower_bound * torch.sigmoid(A_log[:, None].exp() * (g + dt_bias)) actual = kda_gate_reference( g, A_log, dt_bias, safe_gate=True, lower_bound=lower_bound ) torch.testing.assert_close(actual, expected) def test_bwd_gradcheck(): B, T, HV, K = 1, 8, 2, 4 torch.manual_seed(11) g = torch.randn(B, T, HV, K, dtype=torch.float64, requires_grad=True) A_log = torch.randn(HV, dtype=torch.float64, requires_grad=True) * 0.5 dt_bias = torch.randn(HV, K, dtype=torch.float64, requires_grad=True) * 0.1 A_log.requires_grad_(True) dt_bias.requires_grad_(True) def fn(g, A_log, dt_bias): return kda_gate_reference(g, A_log, dt_bias).sum() assert torch.autograd.gradcheck(fn, (g, A_log, dt_bias), eps=1e-6, atol=1e-4) print("L5 bwd-gradcheck: PASSED") if __name__ == "__main__": test_standard_gate_matches_formula() test_safe_gate_matches_formula() test_bwd_gradcheck()