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
+78
View File
@@ -0,0 +1,78 @@
"""Causal-language-model behavior independent of toy memorization."""
import torch
import torch.nn.functional as F
from kda.layers.kda_attn import KDAAttention
from kda.layers.swiglu import SwiGLUMLP
from kda.models.causal_lm import CausalLM
from kda.models.config import KDAConfig
def _model():
torch.manual_seed(51)
config = KDAConfig(
hidden_size=16,
num_hidden_layers=1,
num_heads=2,
num_value_heads=2,
head_dim=4,
chunk_size=4,
vocab_size=32,
intermediate_size=32,
kda_backend="reference",
)
return CausalLM(config).eval()
def test_kda_schedule_and_unified_stem():
config = KDAConfig(num_hidden_layers=2)
assert config.layer_specs() == [("kda", "swiglu"), ("kda", "swiglu")]
model = CausalLM(config)
assert isinstance(model.blocks[0].attn, KDAAttention)
assert isinstance(model.blocks[0].ffn, SwiGLUMLP)
def test_future_token_does_not_change_past_logits():
model = _model()
first = torch.tensor([[1, 2, 3, 4]])
second = torch.tensor([[1, 2, 3, 9]])
with torch.no_grad():
first_logits = model(first)
second_logits = model(second)
torch.testing.assert_close(first_logits[:, :3], second_logits[:, :3])
def test_attention_reads_operator_flags_from_config():
torch.manual_seed(52)
config = KDAConfig(
hidden_size=16,
num_hidden_layers=1,
num_heads=2,
num_value_heads=2,
head_dim=4,
chunk_size=4,
vocab_size=32,
intermediate_size=32,
use_gate_in_kernel=False,
use_qk_l2norm_in_kernel=False,
use_beta_sigmoid_in_kernel=False,
lower_bound=None,
kda_backend="reference",
)
model = CausalLM(config).eval()
with torch.no_grad():
logits = model(torch.tensor([[1, 2, 3, 4]]))
assert logits.shape == (1, 4, 32)
def test_loss_is_shifted_next_token_cross_entropy():
model = _model()
tokens = torch.tensor([[1, 2, 3, 4]])
with torch.no_grad():
logits = model(tokens)
actual = model(tokens, labels=tokens)
expected = F.cross_entropy(
logits[:, :-1].reshape(-1, logits.size(-1)),
tokens[:, 1:].reshape(-1),
)
torch.testing.assert_close(actual, expected)