Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
79 lines
2.2 KiB
Python
79 lines
2.2 KiB
Python
"""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)
|