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
@@ -0,0 +1,64 @@
import torch
import torch.nn.functional as F
from kda.models.causal_lm import CausalLM
from kda.models.config import KDAConfig
from kda.models.k3_config import K3Config
def _tiny():
torch.manual_seed(4)
cfg = KDAConfig(
hidden_size=16,
num_hidden_layers=2,
num_heads=2,
num_value_heads=2,
head_dim=4,
chunk_size=4,
vocab_size=32,
intermediate_size=32,
kda_backend="reference",
)
return CausalLM(cfg), cfg
def test_ignore_index_skips_masked_positions():
model, _ = _tiny()
tokens = torch.tensor([[1, 2, 3, 4]])
labels = tokens.clone()
labels[:, 1:3] = -100
with torch.no_grad():
logits = model(tokens)
actual = model(tokens, labels=labels)
expected = F.cross_entropy(
logits[:, :-1].reshape(-1, logits.size(-1)),
labels[:, 1:].reshape(-1),
ignore_index=-100,
)
torch.testing.assert_close(actual, expected)
def test_gradient_checkpointing_matches_eager_grad():
torch.manual_seed(8)
tokens = torch.randint(0, 32, (2, 8))
m1, cfg = _tiny()
m2 = CausalLM(cfg)
m2.load_state_dict(m1.state_dict())
m2.gradient_checkpointing = True
m1.train()
m2.train()
l1 = m1(tokens, labels=tokens)
l2 = m2(tokens, labels=tokens)
torch.testing.assert_close(l1, l2, atol=1e-5, rtol=1e-5)
l1.backward()
l2.backward()
for p1, p2 in zip(m1.parameters(), m2.parameters()):
if p1.grad is None:
assert p2.grad is None
continue
torch.testing.assert_close(p1.grad, p2.grad, atol=1e-4, rtol=1e-4)
def test_0_5b_preset_enables_checkpointing():
assert K3Config.preset("0.5b").gradient_checkpointing is True
assert K3Config.preset("toy").gradient_checkpointing is False