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,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