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
+1
View File
@@ -0,0 +1 @@
"""Model and training integration tests."""
+176
View File
@@ -0,0 +1,176 @@
"""AttnRes depth mixer: switch, no double-register, causality, two-phase match."""
from dataclasses import asdict
import pytest
import torch
from kda.layers.attn_res import BlockAttnResStack, FullAttnResStack, atomic_block_size
from kda.layers.kda_attn import KDAAttention
from kda.models.causal_lm import CausalLM
from kda.models.config import KDAConfig
from kda.models.k3_config import K3Config
from kda.training.toy import load_ckpt, save_ckpt
def _tiny_k3(**kwargs):
defaults = dict(
hidden_size=32,
num_hidden_layers=4,
num_heads=4,
head_dim=8,
chunk_size=4,
vocab_size=64,
moe_latent_size=16,
moe_d_ff=16,
n_routed=4,
top_k=2,
n_shared=1,
kv_lora_rank=8,
q_lora_rank=16,
qk_nope_head_dim=8,
v_head_dim=8,
)
defaults.update(kwargs)
return K3Config(**defaults)
def test_default_attnres_is_off():
cfg = K3Config(num_hidden_layers=2)
assert cfg.attnres == "off"
model = CausalLM(cfg)
assert model.mixer is None
assert isinstance(model.blocks[0].attn, KDAAttention)
def test_invalid_attnres_rejected():
with pytest.raises(ValueError, match="attnres"):
K3Config(attnres="yes")
with pytest.raises(ValueError, match="attnres_block_size"):
K3Config(attnres="block", attnres_block_size=0)
@pytest.mark.parametrize("mode, stack_cls", [("block", BlockAttnResStack), ("full", FullAttnResStack)])
def test_mixer_kind_and_atomic_count(mode, stack_cls):
cfg = _tiny_k3(attnres=mode, attnres_block_size=2)
model = CausalLM(cfg)
assert model.attnres == mode
assert isinstance(model.mixer, stack_cls)
assert len(model.mixer.layers) == 2 * cfg.num_hidden_layers
if mode == "block":
assert model.mixer.block_size == 4 # 2 DecoderBlocks × attn|ffn
def test_auto_block_size_targets_eight_blocks():
assert atomic_block_size(24, None) == 6 # 3 DecoderBlocks × 2
assert atomic_block_size(4, None) == 2
assert atomic_block_size(93, None) == 24 # 12 DecoderBlocks × 2, K3 S=12
def test_no_duplicate_parameter_ids():
model = CausalLM(_tiny_k3(attnres="block"))
ids = [id(p) for p in model.parameters()]
assert len(ids) == len(set(ids))
names = [n for n, _ in model.named_parameters()]
assert len(names) == len(set(names))
residual_names = [n for n in names if "residuals" in n or "final_residual" in n]
assert residual_names
block_names = [n for n in names if n.startswith("blocks.")]
mixer_weight_names = [
n for n in names if n.startswith("mixer.layers.") and "query" not in n and "norm" not in n
]
assert block_names
assert mixer_weight_names == []
def test_off_and_block_differ_at_same_seed():
torch.manual_seed(0)
off = CausalLM(_tiny_k3(attnres="off"))
torch.manual_seed(0)
on = CausalLM(_tiny_k3(attnres="block"))
x = torch.tensor([[1, 2, 3, 4]])
with torch.no_grad():
assert not torch.allclose(off(x), on(x))
def test_block_two_phase_matches_naive():
torch.manual_seed(4)
model = CausalLM(_tiny_k3(attnres="block", attnres_block_size=2)).eval()
x = torch.randint(0, 64, (2, 8))
with torch.no_grad():
emb = model.embedding(x)
naive = model.mixer.forward_naive(emb)
two_phase = model.mixer(emb)
torch.testing.assert_close(naive, two_phase, atol=1e-5, rtol=1e-5)
@pytest.mark.parametrize("mode", ["block", "full"])
def test_attnres_is_still_causal(mode):
torch.manual_seed(51)
model = CausalLM(_tiny_k3(attnres=mode, num_hidden_layers=2)).eval()
with torch.no_grad():
a = model(torch.tensor([[1, 2, 3, 4]]))
b = model(torch.tensor([[1, 2, 3, 9]]))
torch.testing.assert_close(a[:, :3], b[:, :3], atol=1e-5, rtol=1e-5)
def test_kda_config_block_runs():
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,
attnres="block",
attnres_block_size=1,
)
model = CausalLM(cfg)
logits = model(torch.tensor([[1, 2, 3, 4]]))
assert logits.shape == (1, 4, 32)
def test_attnres_ckpt_roundtrip(tmp_path):
cfg = _tiny_k3(attnres="block", attnres_block_size=2)
model = CausalLM(cfg)
path = str(tmp_path / "attnres.pt")
save_ckpt(model, cfg, path)
loaded, loaded_cfg = load_ckpt(path)
assert loaded_cfg.attnres == "block"
assert loaded_cfg.attnres_block_size == 2
torch.manual_seed(1)
x = torch.randint(0, cfg.vocab_size, (2, 8))
with torch.no_grad():
torch.testing.assert_close(model(x), loaded(x), atol=1e-5, rtol=1e-5)
def test_old_ckpt_without_attnres_stays_off(tmp_path):
cfg = _tiny_k3()
payload = asdict(cfg)
payload.pop("attnres")
payload.pop("attnres_block_size")
payload.pop("attnres_zero_init_queries")
payload.pop("attnres_final_aggregate")
model = CausalLM(cfg)
path = str(tmp_path / "legacy.pt")
torch.save({"model_state": model.state_dict(), "config": payload}, path)
_, loaded_cfg = load_ckpt(path)
assert loaded_cfg.attnres == "off"
assert loaded_cfg.attnres_block_size is None
def test_attnres_block_overfits_single_batch():
torch.manual_seed(30)
cfg = _tiny_k3(attnres="block", num_hidden_layers=2, attnres_block_size=1)
model = CausalLM(cfg)
x = torch.randint(0, cfg.vocab_size, (2, 8))
optim = torch.optim.AdamW(model.parameters(), lr=3e-3)
final = None
for _ in range(200):
optim.zero_grad()
loss = model(x, labels=x)
loss.backward()
optim.step()
final = loss.item()
assert final < 0.5, f"final loss {final:.4f} >= 0.5"
@@ -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
+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)
+56
View File
@@ -0,0 +1,56 @@
"""load_ckpt must read checkpoints written before the ffn/config renames."""
from dataclasses import asdict
import pytest
import torch
from kda.models.config import KDAConfig
from kda.models.k3_config import K3Config
from kda.training.toy import load_ckpt, save_ckpt
def _legacy(state, new, old):
"""Undo the ffn rename, reproducing a pre-rename checkpoint."""
renamed = {
k.replace(f".{new}.", f".{old}.").replace(f".{new}_norm.", f".{old}_norm."): v
for k, v in state.items()
}
assert any(f".{old}." in k for k in renamed), "fixture renamed nothing"
return renamed
@pytest.mark.parametrize(
("config", "old"),
[
(KDAConfig(num_hidden_layers=2), "mlp"), # dense: was named .mlp
(K3Config(num_hidden_layers=2), "moe"), # K3: was named .moe
],
ids=["kda-mlp", "k3-moe"],
)
def test_legacy_ffn_names_still_load(tmp_path, config, old):
from kda.models.causal_lm import CausalLM
model = CausalLM(config)
path = str(tmp_path / "legacy.pt")
torch.save(
{
"model_state": _legacy(model.state_dict(), "ffn", old),
"config": asdict(config),
},
path,
)
loaded, loaded_config = load_ckpt(path)
assert type(loaded_config) is type(config)
for name, want in model.state_dict().items():
torch.testing.assert_close(loaded.state_dict()[name], want)
@pytest.mark.parametrize("config", [KDAConfig(num_hidden_layers=2), K3Config(num_hidden_layers=2)])
def test_roundtrip_picks_the_right_config_class(tmp_path, config):
from kda.models.causal_lm import CausalLM
path = str(tmp_path / "ckpt.pt")
save_ckpt(CausalLM(config), config, path)
_, loaded_config = load_ckpt(path)
assert loaded_config == config
+23
View File
@@ -0,0 +1,23 @@
from pathlib import Path
from kda.training.eval_mt import _instruction
from kda.training.prompts import instruction_prompt
_ROOT = Path(__file__).resolve().parents[2] / "data" / "eval"
def _lines(name: str) -> list[str]:
return [ln.strip() for ln in (_ROOT / name).read_text(encoding="utf-8").splitlines() if ln.strip()]
def test_frozen_eval_files_are_aligned():
zh_src, zh_ref = _lines("zh2en.src.txt"), _lines("zh2en.ref.txt")
en_src, en_ref = _lines("en2zh.src.txt"), _lines("en2zh.ref.txt")
assert len(zh_src) == len(zh_ref) >= 16
assert len(en_src) == len(en_ref) >= 16
assert all("\t" not in s for s in zh_src + en_src)
def test_eval_instruction_is_the_sft_template():
assert _instruction("q", "en") == instruction_prompt("q", "en")
assert _instruction("q", "zh") == instruction_prompt("q", "zh")
+44
View File
@@ -0,0 +1,44 @@
"""translation_success and eval helpers (no GPU, no FLORES download)."""
from kda.training.success import translation_success
def test_empty_and_copy_fail():
src = "人工智能的发展改变了世界。"
assert translation_success(src, "", ref="The development of AI changed the world.", target_lang="en") is False
assert translation_success(src, src, ref="The development of AI changed the world.", target_lang="en") is False
def test_wrong_language_fails():
src = "The cat sat on the mat."
hyp = "The cat sat on the mat and smiled."
ref = "猫坐在垫子上。"
assert translation_success(src, hyp, ref, target_lang="zh") is False
def test_instruction_leak_fails():
src = "Hello"
hyp = "翻译如下:你好"
assert translation_success(src, hyp, ref="你好", target_lang="zh") is False
def test_good_zh2en_passes():
src = "今天天气很好。"
hyp = "The weather is very nice today."
ref = "The weather is very nice today."
assert translation_success(src, hyp, ref, target_lang="en") is True
def test_container_help_exits_2():
import importlib.util
from pathlib import Path
import pytest
path = Path(__file__).resolve().parents[2] / "scripts" / "container_help.py"
spec = importlib.util.spec_from_file_location("container_help", path)
mod = importlib.util.module_from_spec(spec)
assert spec.loader is not None
spec.loader.exec_module(mod)
with pytest.raises(SystemExit) as ei:
mod.main()
assert ei.value.code == 2
+145
View File
@@ -0,0 +1,145 @@
"""K3 架构复现测试: MLA 吸收等价, LatentMoE 路由, hybrid pattern, 因果性, overfit."""
import torch
import torch.nn.functional as F
import pytest
from kda.layers.kda_attn import KDAAttention
from kda.layers.latent_moe import LatentMoE
from kda.layers.mla import GatedMLA
from kda.models.causal_lm import CausalLM
from kda.models.k3_config import K3Config
def _mla(d=64, H=4, r=16, q_r=32, d_q=16, d_v=16):
torch.manual_seed(7)
return GatedMLA(d, H, r, q_r, d_q, d_v)
def _naive_mla(x, module: GatedMLA):
"""解压版参考: 标准 attention (吸收版数学上应与它逐位一致)."""
B, T, _ = x.shape
H, r = module.num_heads, module.kv_up.in_features
c = module.kv_norm(module.kv_down(x)) # [B,T,r]
q = module.q_up(module.q_norm(module.q_down(x))).view(B, T, H, module.qk_nope_head_dim)
w = module.kv_up.weight
w_uk = w[: H * module.qk_nope_head_dim].view(H, module.qk_nope_head_dim, r)
w_uv = w[H * module.qk_nope_head_dim :].view(H, module.v_head_dim, r)
k = torch.einsum("btj,hvj->bthv", c, w_uk) # 解压 K
v = torch.einsum("btj,hvj->bthv", c, w_uv) # 解压 V
scores = torch.einsum("bthv,bshv->bhts", q, k) # [B,H,T,T]
mask = torch.triu(torch.ones(T, T, dtype=torch.bool), diagonal=1)
scores = scores.masked_fill(mask, float("-inf"))
attn = F.softmax(scores, dim=-1)
o = torch.einsum("bhts,bshv->bthv", attn, v) # [B,T,H,d_v]
o = o.reshape(B, T, H * module.v_head_dim)
gate = torch.sigmoid(module.gate(x))
return module.o_proj(gate * o)
def test_mla_absorption_matches_unrolled():
m = _mla().eval()
x = torch.randn(3, 12, 64)
with torch.no_grad():
absorbed = m(x)
unrolled = _naive_mla(x, m)
torch.testing.assert_close(absorbed, unrolled, atol=1e-5, rtol=1e-5)
def test_mla_absorption_matches_unrolled_grad():
"""吸收版与解压版的梯度也应一致 (fwd+bwd 双重验证)."""
m1, m2 = _mla(), _mla()
m2.load_state_dict(m1.state_dict())
x = torch.randn(2, 8, 64)
l1 = m1(x).square().mean()
l2 = _naive_mla(x, m2).square().mean()
l1.backward()
l2.backward()
for (n1, p1), (n2, p2) in zip(m1.named_parameters(), m2.named_parameters()):
torch.testing.assert_close(p1.grad, p2.grad, atol=1e-5, rtol=1e-5)
def test_hybrid_layer_pattern():
cfg = K3Config(num_hidden_layers=4)
assert cfg.layer_types() == ["kda", "kda", "kda", "mla"]
cfg8 = K3Config(num_hidden_layers=8)
assert cfg8.layer_types() == ["kda", "kda", "kda", "mla"] * 2
# 末层强制 MLA: L=5 → 层 3 MLA + 层 4 (末层) MLA
cfg5 = K3Config(num_hidden_layers=5)
assert cfg5.layer_types() == ["kda", "kda", "kda", "mla", "mla"]
assert cfg.layer_specs() == [("kda", "moe"), ("kda", "moe"), ("kda", "moe"), ("mla", "moe")]
model = CausalLM(K3Config(num_hidden_layers=4, hidden_size=32, moe_d_ff=16, moe_latent_size=16))
assert isinstance(model.blocks[0].attn, KDAAttention)
assert isinstance(model.blocks[3].attn, GatedMLA)
assert isinstance(model.blocks[0].ffn, LatentMoE)
def test_preset_0_5b_schedule():
cfg = K3Config.preset("0.5b")
assert cfg.hidden_size == 768
assert cfg.num_heads * cfg.head_dim == cfg.hidden_size
assert cfg.num_hidden_layers == 24
assert cfg.tie_word_embeddings
assert cfg.chunk_size == 64
assert cfg.gradient_checkpointing is True
assert cfg.moe_latent_size == cfg.hidden_size // 2
types = cfg.layer_types()
assert types.count("mla") == 6
assert types[-1] == "mla"
assert cfg.layer_specs()[3] == ("mla", "moe")
def test_moe_router_activates_topk_only():
from kda.layers.latent_moe import LatentMoE
torch.manual_seed(3)
moe = LatentMoE(hidden_size=32, latent_size=16, n_routed=8, top_k=2, n_shared=1, d_ff=24)
x = torch.randn(2, 6, 32)
with torch.no_grad():
y = moe(x)
logits = moe.router(x)
topk = torch.topk(logits, moe.top_k, dim=-1)
z = moe.down(x)
# 手算: 只有 top-k 专家输出被加权, 再经 shared + up(norm(u))
expected_u = torch.zeros(2, 6, moe.latent_size)
all_out = torch.stack([e(z) for e in moe.experts]) # [R,B,T,ℓ]
probs = F.softmax(topk.values, dim=-1)
for i in range(moe.top_k):
idx = topk.indices[:, :, i]
for b in range(2):
for t in range(6):
expected_u[b, t] += probs[b, t, i] * all_out[idx[b, t], b, t]
shared = torch.stack([e(x) for e in moe.shared]).sum(0)
expected_y = shared + moe.up(moe.norm(expected_u))
torch.testing.assert_close(y, expected_y, atol=1e-5, rtol=1e-5)
assert moe.last_route_ids is not None
assert moe.last_route_ids.shape[-1] == moe.top_k
def test_k3_causal_future_does_not_change_past_logits():
torch.manual_seed(51)
cfg = K3Config(hidden_size=64, num_hidden_layers=4, num_heads=4, head_dim=8,
chunk_size=4, vocab_size=64, moe_latent_size=32, moe_d_ff=24,
n_routed=8, kv_lora_rank=16, q_lora_rank=32, qk_nope_head_dim=8, v_head_dim=8)
m = CausalLM(cfg).eval()
with torch.no_grad():
a = m(torch.tensor([[1, 2, 3, 4]]))
b = m(torch.tensor([[1, 2, 3, 9]]))
torch.testing.assert_close(a[:, :3], b[:, :3], atol=1e-6, rtol=0)
def test_k3_small_model_overfits_single_batch():
"""K3 混合架构单 batch overfit 冒烟: loss < 0.5 (收敛即架构可训)."""
torch.manual_seed(30)
cfg = K3Config(hidden_size=64, num_hidden_layers=2, num_heads=4, head_dim=8,
chunk_size=4, vocab_size=64, moe_latent_size=32, moe_d_ff=24,
n_routed=8, kv_lora_rank=16, q_lora_rank=32, qk_nope_head_dim=8, v_head_dim=8)
m = CausalLM(cfg)
x = torch.randint(0, cfg.vocab_size, (2, 16))
optim = torch.optim.AdamW(m.parameters(), lr=3e-3)
final = None
for step in range(200):
optim.zero_grad()
loss = m(x, labels=x)
loss.backward()
optim.step()
final = loss.item()
assert final < 0.5, f"final loss {final:.4f} >= 0.5"
+53
View File
@@ -0,0 +1,53 @@
import json
import sys
import torch
from kda.training.data import (
chunk_ids,
fetch_wiki_texts,
interleave_balanced,
split_heldout,
)
def test_interleave_is_one_to_one_monolingual_blocks():
a = list(range(10))
b = list(range(100, 112))
out = interleave_balanced(a, b, block=4)
assert out == [0, 1, 2, 3, 100, 101, 102, 103, 4, 5, 6, 7, 104, 105, 106, 107]
def test_split_heldout_keeps_at_least_one_train():
chunks = torch.arange(10).view(10, 1, 1)
train, held = split_heldout(chunks, frac=0.01, min_heldout=1)
assert train.size(0) == 9
assert held.size(0) == 1
empty_train, empty_held = split_heldout(chunks[:1], frac=0.5)
assert empty_train.size(0) == 1
assert empty_held.size(0) == 0
def test_chunk_ids_drops_tail():
ids = list(range(10))
chunks = chunk_ids(ids, batch=2, seq_len=4)
assert chunks.shape == (1, 2, 4)
def test_wiki_cache_roundtrip(tmp_path, monkeypatch):
cache = tmp_path / "pretrain"
cache.mkdir()
path = cache / "wiki-zh-n2-limit3.jsonl"
path.write_text(
"\n".join(json.dumps({"text": f"article {i}"}) for i in range(3)) + "\n",
encoding="utf-8",
)
def _boom(*_a, **_k):
raise AssertionError("must not hit the network")
fake = type(sys)("datasets")
fake.load_dataset = _boom
monkeypatch.setitem(sys.modules, "datasets", fake)
texts = fetch_wiki_texts(3, lang="zh", cache_dir=cache)
assert texts == ["article 0", "article 1", "article 2"]
+21
View File
@@ -0,0 +1,21 @@
from kda.training.schedule import lr_scale, total_opt_steps
def test_warmup_then_cosine_floor():
assert abs(lr_scale(0, warmup=10, total_opt=100) - 0.1) < 1e-9
assert abs(lr_scale(9, warmup=10, total_opt=100) - 1.0) < 1e-9
assert abs(lr_scale(10, warmup=10, total_opt=100) - 1.0) < 1e-6
end = lr_scale(99, warmup=10, total_opt=100)
assert abs(end - 0.1) < 1e-6
def test_horizon_prefers_the_earlier_stop():
# 8.2M tokens @ batch 2 seq 2048 acc 8 -> 250 opt
opt_from_tokens = total_opt_steps(
max_tokens=8_192_000, max_micro=10_000, batch=2, seq_len=2048, grad_acc=8
)
assert opt_from_tokens == 250
opt_from_micro = total_opt_steps(
max_tokens=10**12, max_micro=2000, batch=2, seq_len=2048, grad_acc=8
)
assert opt_from_micro == 250
+52
View File
@@ -0,0 +1,52 @@
from kda.training.data import IGNORE_INDEX, collate_sft, encode_sft_row, load_sft_rows
from kda.training.prompts import instruction_prompt
class _Tok:
vocab_size = 32
def encode(self, text: str) -> list[int]:
return [min((ord(c) % 30) + 1, 31) for c in text[:12]] or [1]
def decode(self, ids: list[int]) -> str:
return "x" * len(ids)
def test_instruction_matches_eval_template():
assert instruction_prompt("你好", "en") == "Translate to English:\n你好"
assert instruction_prompt("Hello", "zh") == "Translate to Chinese:\nHello"
def test_prompt_tokens_are_ignored():
tok = _Tok()
src, tgt = "ab", "cd"
ids, labels = encode_sft_row(tok, src, tgt, "en", max_len=64)
prompt_n = len(tok.encode(instruction_prompt(src, "en")))
assert labels[:prompt_n] == [IGNORE_INDEX] * prompt_n
assert all(v != IGNORE_INDEX for v in labels[prompt_n:])
assert ids[prompt_n:] == tok.encode(tgt)
def test_collate_and_jsonl(tmp_path):
path = tmp_path / "tiny.jsonl"
path.write_text(
'{"src": "a", "tgt": "b", "target_lang": "en"}\n'
'{"src": "c", "tgt": "d", "target_lang": "zh"}\n',
encoding="utf-8",
)
rows = load_sft_rows(path)
assert len(rows) == 2
x, y = collate_sft(rows, _Tok(), max_len=32)
assert x.shape == y.shape
assert x.size(0) == 2
assert (y == IGNORE_INDEX).any()
def test_toy_sft_file_parses():
from pathlib import Path
path = Path(__file__).resolve().parents[2] / "data" / "sft" / "toy.jsonl"
rows = load_sft_rows(path)
assert len(rows) >= 20
langs = {r["target_lang"] for r in rows}
assert langs == {"en", "zh"}
+193
View File
@@ -0,0 +1,193 @@
"""TensorLens 集成测试: KDA 模型张量 -> trace -> 全局 store -> Flask 端点.
依赖: tensorlens (未安装时整个模块 skip, 用 `pip install tensorlens` 启用).
测三层:
1. trace — KDA 前向的真实 1D/2D/3D 张量 (embed/block 输出/logits/权重)
规范化成 int8 存入 tensorlens 全局 store
2. normalize — 四种策略 (clip/minmax/zscore/none) 的边界行为
3. HTTP — 用 Flask test_client 验证 /api/list_tensors 与 /api/get_tensor,
不启动阻塞的 gunicorn server (viewer() 为交互式入口, 不做自动化)
"""
import numpy as np
import pytest
import torch
pytest.importorskip("tensorlens")
from tensorlens.core import global_store
from tensorlens.tensorlens import normalize_to_int8, trace
from tensorlens.web.server import app
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()
@pytest.fixture(autouse=True)
def _clean_store():
"""global_store 是模块级单例, 每个测试前后清空, 避免互相污染."""
global_store.INMEMORY_TENSORS.clear()
yield
global_store.INMEMORY_TENSORS.clear()
def _trace_kda_tensors(model):
"""前向一次, 把 KDA 的 1D/2D/3D 张量全部 trace 进 store, 返回 logits numpy."""
x = torch.tensor([[1, 2, 3, 4]])
hidden = {}
def hook_fn(name):
def hook(module, inp, out):
hidden[name] = out.detach()
return hook
model.embedding.register_forward_hook(hook_fn("embed"))
model.blocks[0].register_forward_hook(hook_fn("block0"))
with torch.no_grad():
logits = model(x)
logits_np = logits.detach().numpy() # [1, T, vocab] 3D
trace("lm_head.weight", model.lm_head.weight.detach().numpy()) # [vocab, hidden] 2D
trace("embed", hidden["embed"].numpy()) # [1, T, hidden] 3D
trace("block0.out", hidden["block0"].numpy()) # [1, T, hidden] 3D
trace("logits", logits_np) # [1, T, vocab] 3D
trace("logits.row0", logits_np[0, 0]) # [vocab] 1D
return logits_np
# ---------------------------------------------------------------------------
# Layer 1: trace — KDA 张量进 store
# ---------------------------------------------------------------------------
def test_trace_stores_int8_with_expected_shape():
model = _model()
logits_np = _trace_kda_tensors(model)
store = global_store.INMEMORY_TENSORS
assert set(store) == {"lm_head.weight", "embed", "block0.out", "logits", "logits.row0"}
assert store["logits"].dtype == np.int8
assert store["logits"].shape == logits_np.shape
assert store["lm_head.weight"].shape == (32, 16)
assert store["logits.row0"].ndim == 1
assert store["logits"].min() >= -128 and store["logits"].max() <= 127
# ---------------------------------------------------------------------------
# Layer 2: normalize_to_int8 — 四种策略边界行为
# ---------------------------------------------------------------------------
def test_clip_normalization_bounds():
t = np.array([[-100.0, 0.0, 100.0]])
out = normalize_to_int8(t, (-1.0, 1.0), "clip")
assert out.dtype == np.int8
assert out[0, 0] == -127 and out[0, 2] == 127
assert out[0, 1] == 0
def test_minmax_constant_tensor_returns_zeros():
t = np.full((2, 3), 0.5)
out = normalize_to_int8(t, (-1.0, 1.0), "minmax")
assert (out == 0).all()
def test_zscore_zero_std_returns_zeros():
t = np.ones((4, 4))
out = normalize_to_int8(t, (-1.0, 1.0), "zscore")
assert (out == 0).all()
def test_none_strategy_scales_by_127():
t = np.array([[0.5, -0.5]])
out = normalize_to_int8(t, (-1.0, 1.0), "none")
assert out[0, 0] == 63 and out[0, 1] == -63 # int8 cast 向零截断
def test_unsupported_normalization_raises():
with pytest.raises(ValueError):
normalize_to_int8(np.zeros(3), (-1.0, 1.0), "bogus")
# ---------------------------------------------------------------------------
# Layer 2.5: trace 输入校验
# ---------------------------------------------------------------------------
def test_trace_rejects_non_ndarray():
with pytest.raises(TypeError):
trace("bad", torch.zeros(3))
def test_trace_rejects_empty_key():
with pytest.raises(ValueError):
trace("", np.zeros(3))
# ---------------------------------------------------------------------------
# Layer 3: Flask 端点 (test_client, 不起真实 server)
# ---------------------------------------------------------------------------
def test_list_tensors_endpoint_reports_kda_tensors():
model = _model()
_trace_kda_tensors(model)
resp = app.test_client().get("/api/list_tensors")
assert resp.status_code == 200
body = resp.get_json()
assert body["count"] == 5
keys = [t["key"] for t in body["available_tensors"]]
assert "logits" in keys and "lm_head.weight" in keys
def test_get_tensor_endpoint_returns_data():
model = _model()
logits_np = _trace_kda_tensors(model)
resp = app.test_client().get("/api/get_tensor?tensor_key=logits")
assert resp.status_code == 200
body = resp.get_json()
assert body["shape"] == list(logits_np.shape)
assert len(body["data"]) == logits_np.shape[0]
def test_get_tensor_missing_key_returns_400():
model = _model()
_trace_kda_tensors(model)
resp = app.test_client().get("/api/get_tensor")
assert resp.status_code == 400
assert "tensor_key" in resp.get_json()["error"]
def test_get_tensor_unknown_key_returns_404():
resp = app.test_client().get("/api/get_tensor?tensor_key=nope")
assert resp.status_code == 404
def test_config_endpoint():
resp = app.test_client().get("/api/config")
assert resp.status_code == 200
assert resp.get_json()["status"] == "ok"
if __name__ == "__main__":
import sys
sys.exit(pytest.main([__file__, "-v"]))
+83
View File
@@ -0,0 +1,83 @@
"""TorchLens 集成测试: 计算图展开 KDA 模型并提取逐层激活.
依赖: torchlens (未安装时整个模块 skip, 用 `pip install torchlens` 启用).
验证三点:
1. trace 能展开 KDA 模型 —— 关键子模块 (embedding / attention 各投影 /
block 输出 / norm / lm_head 输出) 的激活被捕获且形状正确
2. 展开不改变模型行为 —— trace 记录的输出与直接 forward 严格一致
3. extract 便捷接口 —— 按模块名批量取激活
torchlens 2.34 的 trace[key] 返回 Op 对象, 取原始 tensor 用 `.tensor`.
"""
import pytest
import torch
pytest.importorskip("torchlens")
import torchlens as tl
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_trace_captures_kda_submodule_activations():
model = _model()
x = torch.tensor([[1, 2, 3, 4]])
with torch.no_grad():
trace = tl.trace(model, x, capture=tl.options.CaptureOptions(verbose=False))
# 模块激活被捕获, 形状正确
assert tuple(trace["embedding"].tensor.shape) == (1, 4, 16)
assert tuple(trace["blocks.0.attn.q_proj"].tensor.shape) == (1, 4, 8) # H*K = 2*4
assert tuple(trace["blocks.0.attn.v_proj"].tensor.shape) == (1, 4, 8) # HV*V = 2*4
assert tuple(trace["blocks.0.attn"].tensor.shape) == (1, 4, 16)
assert tuple(trace["blocks.0.ffn"].tensor.shape) == (1, 4, 16)
assert tuple(trace["norm"].tensor.shape) == (1, 4, 16)
assert tuple(trace["output"].tensor.shape) == (1, 4, 32) # vocab_size
def test_trace_does_not_change_model_behavior():
model = _model()
x = torch.tensor([[1, 2, 3, 4]])
with torch.no_grad():
trace = tl.trace(model, x, capture=tl.options.CaptureOptions(verbose=False))
traced_logits = trace["output"].tensor
direct_logits = model(x)
torch.testing.assert_close(traced_logits, direct_logits)
def test_extract_returns_activations_by_module_name():
model = _model()
x = torch.tensor([[1, 2, 3, 4]])
with torch.no_grad():
acts = tl.extract(
model, x, ["embedding", "blocks.0.attn.q_proj", "blocks.0", "output"]
)
assert set(acts) == {"embedding", "blocks.0.attn.q_proj", "blocks.0", "output"}
assert tuple(acts["embedding"].shape) == (1, 4, 16)
assert tuple(acts["blocks.0.attn.q_proj"].shape) == (1, 4, 8)
assert tuple(acts["blocks.0"].shape) == (1, 4, 16)
assert tuple(acts["output"].shape) == (1, 4, 32)
if __name__ == "__main__":
import sys
sys.exit(pytest.main([__file__, "-v"]))
+32
View File
@@ -0,0 +1,32 @@
"""L7: toy overfit smoke test. 320 steps loss < 0.1."""
import torch
from kda.models.causal_lm import CausalLM
from kda.models.config import KDAConfig
def test_overfit():
cfg = KDAConfig() # 起步默认 toy 配置
torch.manual_seed(30)
model = CausalLM(cfg).cuda()
x = torch.randint(0, cfg.vocab_size, (4, 16), device="cuda")
labels = x.clone()
# A single repeated batch is an optimizer/dataflow smoke test, so converge it quickly.
optim = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.01)
for step in range(320):
optim.zero_grad()
loss = model(x, labels=labels)
loss.backward()
optim.step()
if step % 64 == 0 or step == 319:
print(f" step {step:3d} loss {loss.item():.4f}")
final = loss.item()
assert final < 0.1, f"final loss {final:.4f} > 0.1"
print(f"L7 overfit: PASSED (final loss {final:.4f})")
if __name__ == "__main__":
test_overfit()