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,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)
|
||||
Reference in New Issue
Block a user