Files
K3/tests/correctness/test_gate.py
T
dela 584f7e9e73 Initial K3 snapshot: 0.5B KDA/MLA/MoE train path
Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias
backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
2026-08-25 14:43:17 +08:00

53 lines
1.8 KiB
Python

"""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()