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.
This commit is contained in:
dela
2026-08-25 14:43:17 +08:00
commit 584f7e9e73
140 changed files with 21592 additions and 0 deletions
+30
View File
@@ -0,0 +1,30 @@
"""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()