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:
@@ -0,0 +1,20 @@
|
||||
"""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)
|
||||
Reference in New Issue
Block a user