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 @@
|
||||
"""Local kernel implementation tests."""
|
||||
@@ -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)
|
||||
@@ -0,0 +1,67 @@
|
||||
"""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()
|
||||
@@ -0,0 +1,29 @@
|
||||
"""L3: vendored FLA Triton fwd vs naive chunked."""
|
||||
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_triton_fwd_matches_naive():
|
||||
B, T, H, HV, K, V = 2, 64, 4, 8, 32, 32
|
||||
torch.manual_seed(3)
|
||||
device = "cuda"
|
||||
q = F.normalize(torch.randn(B, T, H, K, device=device), dim=-1)
|
||||
k = F.normalize(torch.randn(B, T, H, K, device=device), dim=-1)
|
||||
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_ref, S_ref = naive_chunk_kda(
|
||||
q, k, v, g, beta, output_final_state=True, chunk_size=64
|
||||
)
|
||||
o_tr, S_tr = chunk_kda_fwd(
|
||||
q, k, v, g, beta, output_final_state=True, chunk_size=64
|
||||
)
|
||||
torch.testing.assert_close(o_tr.float(), o_ref.float(), rtol=2e-3, atol=2e-3)
|
||||
torch.testing.assert_close(S_tr.float(), S_ref.float(), rtol=2e-3, atol=2e-3)
|
||||
Reference in New Issue
Block a user