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 @@
# allow running tests from package root
+1
View File
@@ -0,0 +1 @@
"""Reference correctness tests."""
+107
View File
@@ -0,0 +1,107 @@
"""Public backend selection must be explicit and reproducible."""
import pytest
import torch
from kda.models.config import KDAConfig
from kda.ops.api import chunk_kda
def _inputs():
torch.manual_seed(41)
B, T, H, HV, K, V = 1, 4, 1, 2, 2, 2
return (
torch.randn(B, T, H, K),
torch.randn(B, T, H, K),
torch.randn(B, T, HV, V),
-torch.rand(B, T, HV, K),
torch.rand(B, T, HV),
)
def test_default_config_enables_kernel_side_transforms():
config = KDAConfig()
assert config.use_gate_in_kernel
assert config.use_qk_l2norm_in_kernel
assert config.use_beta_sigmoid_in_kernel
assert config.lower_bound == -5.0
def test_default_backend_is_local_reference():
inputs = _inputs()
default_output, _ = chunk_kda(*inputs, chunk_size=4, use_qk_l2norm_in_kernel=True)
reference_output, _ = chunk_kda(*inputs, chunk_size=4, backend="reference", use_qk_l2norm_in_kernel=True)
torch.testing.assert_close(default_output, reference_output)
assert KDAConfig().kda_backend == "reference"
def test_auto_is_deprecated_reference_alias():
inputs = _inputs()
with pytest.warns(DeprecationWarning, match="selects the local reference backend"):
auto_output, _ = chunk_kda(*inputs, chunk_size=4, backend="auto", use_qk_l2norm_in_kernel=True)
reference_output, _ = chunk_kda(*inputs, chunk_size=4, backend="reference", use_qk_l2norm_in_kernel=True)
torch.testing.assert_close(auto_output, reference_output)
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
def test_triton_backend_matches_reference():
torch.manual_seed(41)
B, T, H, HV, K, V = 2, 64, 4, 8, 32, 32
device = "cuda"
q = torch.randn(B, T, H, K, device=device)
k = torch.randn(B, T, H, K, device=device)
v = torch.randn(B, T, HV, V, device=device)
g = -torch.rand(B, T, HV, K, device=device)
beta = torch.rand(B, T, HV, device=device)
triton_out, _ = chunk_kda(
q, k, v, g, beta, chunk_size=64, backend="triton", use_qk_l2norm_in_kernel=True
)
reference_out, _ = chunk_kda(
q, k, v, g, beta, chunk_size=64, backend="reference", use_qk_l2norm_in_kernel=True
)
torch.testing.assert_close(
triton_out.float(), reference_out.float(), atol=2e-3, rtol=2e-3
)
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
def test_unnormalized_parity_is_not_meaningful():
"""对拍必须用与训练一致的 q/k L2-norm (守卫测试).
M = I + tril(A_kk*beta, -1) 是单位下三角, 行列式恒为 1、特征值全为 1,
不存在"失去对角占优"一说. 真正的机制是 M = I + N 中 N 严格下三角 (幂零),
故 M^-1 = sum_{m<C} (-N)^m, ‖M^-1‖ ~ ‖N‖^C —— 关于 chunk 长度指数增长.
L2-norm 的作用是把 ‖N‖ 压到 1 以下让这个 Neumann 级数收敛.
不归一化时 ‖k‖>1, 精确解本身就以 ‖k‖^chunk_size 增长 (实测输出 1e31 量级
乃至 inf). 这不是某个后端的缺陷: fp64 逐步递推 (完全不含三角求解) 与分块
形式吻合到 8e-15, 说明爆炸就是该递推的真值. 两个后端的相对误差其实很接近
(1.3e-3 → 2.6e-2), 但绝对差随输出量级一起走, 任何一致性容差都必然失败.
"""
torch.manual_seed(41)
B, T, H, HV, K, V = 2, 64, 4, 8, 32, 32
device = "cuda"
q = torch.randn(B, T, H, K, device=device) # 故意不归一化
k = torch.randn(B, T, H, K, device=device)
v = torch.randn(B, T, HV, V, device=device)
g = -torch.rand(B, T, HV, K, device=device)
beta = torch.rand(B, T, HV, device=device)
with pytest.warns(RuntimeWarning, match="use_qk_l2norm_in_kernel=False"):
triton_out, _ = chunk_kda(q, k, v, g, beta, chunk_size=64, backend="triton")
reference_out, _ = chunk_kda(q, k, v, g, beta, chunk_size=64, backend="reference")
diff = (triton_out.float() - reference_out.float()).abs().max()
assert not torch.isfinite(diff).item() or diff.item() > 1e-2, (
f"unnormalised parity diff {diff.item():.3e} should be pathological"
)
def test_triton_backend_does_not_import_upstream_fla():
import sys
for name in list(sys.modules):
if name == "fla" or name.startswith("fla."):
sys.modules.pop(name)
from kda.ops.triton.chunk import chunk_kda as _triton_chunk_kda # noqa: F401
assert not any(n == "fla" or n.startswith("fla.") for n in sys.modules)
+47
View File
@@ -0,0 +1,47 @@
"""L2: chunked naive vs L1 naive (fwd) + gradcheck."""
import torch
from kda.ops.reference.chunkwise import naive_chunk_kda
from kda.ops.reference.recurrent import naive_kda
def test_fwd_matches_naive():
"""chunked vs naive fwd, atol=1e-4."""
B, T, H, HV, K, V = 2, 32, 2, 4, 8, 8
torch.manual_seed(1)
q = torch.randn(B, T, H, K, dtype=torch.float64)
k = torch.randn(B, T, H, K, dtype=torch.float64)
v = torch.randn(B, T, HV, V, dtype=torch.float64)
g = torch.randn(B, T, HV, K, dtype=torch.float64) * 0.1
beta = torch.rand(B, T, HV, dtype=torch.float64)
o_ref, _ = naive_kda(q, k, v, g, beta, output_final_state=True)
o_chk, _ = naive_chunk_kda(q, k, v, g, beta,
output_final_state=True, chunk_size=8)
diff = (o_ref - o_chk).abs().max().item()
assert diff < 1e-4, f"chunk vs naive fwd max diff {diff:.2e} > 1e-4"
print(f"L2 fwd-vs-naive: PASSED (max diff {diff:.2e})")
def test_gradcheck():
"""L2 chunked gradcheck atol=1e-4 (allowing rtol 1e-3 for triangular path)."""
B, T, H, HV, K, V = 2, 16, 2, 4, 4, 4
torch.manual_seed(2)
q = torch.randn(B, T, H, K, dtype=torch.float64, requires_grad=True)
k = torch.randn(B, T, H, K, dtype=torch.float64, requires_grad=True)
v = torch.randn(B, T, HV, V, dtype=torch.float64, requires_grad=True)
g = torch.randn(B, T, HV, K, dtype=torch.float64, requires_grad=True) * 0.1
beta = torch.rand(B, T, HV, dtype=torch.float64, requires_grad=True)
assert torch.autograd.gradcheck(
lambda q, k, v, g, b: naive_chunk_kda(q, k, v, g, b, chunk_size=4)[0],
(q, k, v, g, beta),
eps=1e-6, atol=1e-4, rtol=1e-3,
), "L2 gradcheck 失败"
print("L2 gradcheck: PASSED")
if __name__ == "__main__":
test_fwd_matches_naive()
test_gradcheck()
+99
View File
@@ -0,0 +1,99 @@
"""Saturated gates must not overflow the reference ``exp``.
The existing kernel tests draw ``g = -rand(...)``, which keeps the per-chunk
gate span near 1 per step and never exercises the exponent budget. A trained
``safe_gate`` model pins whole channels at ``|g| = |lower_bound|`` for a full
chunk, which used to overflow ``_decayed_dot`` and produce NaN for any
``chunk_size > 16``.
"""
import pytest
import torch
from kda.ops.api import chunk_kda
from kda.ops.reference.chunkwise import DECAY_BLOCK, _EXP_LIMIT
LOWER_BOUND = -5.0
def _saturated_inputs(T, device="cpu", seed=0):
"""Gates pinned at ``lower_bound`` on one channel, mild elsewhere."""
torch.manual_seed(seed)
B, H, K, V = 1, 2, 8, 8
q = torch.randn(B, T, H, K, device=device)
k = torch.randn(B, T, H, K, device=device)
v = torch.randn(B, T, H, V, device=device)
g = -torch.rand(B, T, H, K, device=device) * 0.1
g[..., 0] = LOWER_BOUND # fully saturated channel, the worst case
beta = torch.rand(B, T, H, device=device)
return q, k, v, g, beta
def _run(inputs, **kwargs):
kwargs.setdefault("use_qk_l2norm_in_kernel", True)
return chunk_kda(*inputs, **kwargs)[0]
@pytest.mark.parametrize("chunk_size", [16, 32, 64])
def test_saturated_gate_does_not_overflow(chunk_size):
out = _run(_saturated_inputs(128), chunk_size=chunk_size)
assert torch.isfinite(out).all(), f"NaN/inf at chunk_size={chunk_size}"
def test_chunk_size_does_not_change_the_result():
"""Chunking is exact algebra, so every chunk_size must agree."""
inputs = _saturated_inputs(128)
base = _run(inputs, chunk_size=16)
for chunk_size in (32, 64):
other = _run(inputs, chunk_size=chunk_size)
torch.testing.assert_close(other, base, atol=1e-5, rtol=1e-5)
def test_gradients_flow_through_the_row_blocks():
"""_decayed_dot writes its row blocks into a preallocated buffer."""
q, k, v, g, beta = _saturated_inputs(64)
for t in (q, k, v, g, beta):
t.requires_grad_(True)
_run((q, k, v, g, beta), chunk_size=64).sum().backward()
for name, t in zip("qkvgb", (q, k, v, g, beta)):
assert t.grad is not None and torch.isfinite(t.grad).all(), name
def test_decay_block_fits_the_exp_budget():
assert DECAY_BLOCK * abs(LOWER_BOUND) < _EXP_LIMIT
def test_lower_bound_beyond_the_budget_is_rejected():
too_deep = -(_EXP_LIMIT / DECAY_BLOCK) - 1.0
with pytest.raises(ValueError, match="overflows exp"):
_run(
_saturated_inputs(32),
chunk_size=32,
safe_gate=True,
lower_bound=too_deep,
use_gate_in_kernel=True,
A_log=torch.zeros(2),
)
@pytest.mark.parametrize("chunk_size", [32, 50, 64])
def test_unbounded_gate_past_the_budget_warns(chunk_size):
"""safe_gate is bounded, but ``-A.exp() * softplus(x)`` is not."""
q, k, v, g, beta = _saturated_inputs(100)
g[...] = -(_EXP_LIMIT / DECAY_BLOCK) - 0.5 # just over the per-block budget
with pytest.warns(RuntimeWarning, match="gate span within a"):
_run((q, k, v, g, beta), chunk_size=chunk_size)
def test_unnormalised_qk_warns():
with pytest.warns(RuntimeWarning, match="Neumann series"):
_run(_saturated_inputs(32), chunk_size=32, use_qk_l2norm_in_kernel=False)
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
def test_saturated_gate_matches_triton():
inputs = _saturated_inputs(128, device="cuda")
reference = _run(inputs, chunk_size=64, backend="reference")
triton = _run(inputs, chunk_size=64, backend="triton")
torch.testing.assert_close(
triton.float(), reference.float(), atol=2e-3, rtol=2e-3
)
+52
View File
@@ -0,0 +1,52 @@
"""KDA gate reference semantics and autograd coverage."""
import torch
import torch.nn.functional as F
from kda.ops.reference.gate import kda_gate_reference
def test_standard_gate_matches_formula():
B, T, HV, K = 1, 32, 4, 8
torch.manual_seed(10)
g = torch.randn(B, T, HV, K, dtype=torch.float64)
A_log = torch.randn(HV, dtype=torch.float64) * 0.5
dt_bias = torch.randn(HV, K, dtype=torch.float64) * 0.1
expected = -A_log[:, None].exp() * F.softplus(g + dt_bias)
actual = kda_gate_reference(g, A_log, dt_bias)
torch.testing.assert_close(actual, expected)
def test_safe_gate_matches_formula():
B, T, HV, K = 1, 8, 2, 4
torch.manual_seed(12)
g = torch.randn(B, T, HV, K, dtype=torch.float64)
A_log = torch.randn(HV, dtype=torch.float64)
dt_bias = torch.randn(HV, K, dtype=torch.float64)
lower_bound = -5.0
expected = lower_bound * torch.sigmoid(A_log[:, None].exp() * (g + dt_bias))
actual = kda_gate_reference(
g, A_log, dt_bias, safe_gate=True, lower_bound=lower_bound
)
torch.testing.assert_close(actual, expected)
def test_bwd_gradcheck():
B, T, HV, K = 1, 8, 2, 4
torch.manual_seed(11)
g = torch.randn(B, T, HV, K, dtype=torch.float64, requires_grad=True)
A_log = torch.randn(HV, dtype=torch.float64, requires_grad=True) * 0.5
dt_bias = torch.randn(HV, K, dtype=torch.float64, requires_grad=True) * 0.1
A_log.requires_grad_(True)
dt_bias.requires_grad_(True)
def fn(g, A_log, dt_bias):
return kda_gate_reference(g, A_log, dt_bias).sum()
assert torch.autograd.gradcheck(fn, (g, A_log, dt_bias), eps=1e-6, atol=1e-4)
print("L5 bwd-gradcheck: PASSED")
if __name__ == "__main__":
test_standard_gate_matches_formula()
test_safe_gate_matches_formula()
test_bwd_gradcheck()
+34
View File
@@ -0,0 +1,34 @@
"""L1: gradcheck for naive_recurrent_kda.
验证策略: torch.autograd.gradcheck 走 forward+backward 五个梯度.
强制 dtype=float64; eps=1e-6, atol=1e-4.
shape (small):
B=2, T=8, H=2, HV=4, K=4, V=4
"""
import torch
from kda.ops.reference.recurrent import naive_kda
def test_gradcheck():
B, T, H, HV, K, V = 2, 8, 2, 4, 4, 4
torch.manual_seed(0)
# 所有输入都需要 requires_grad=True
q = torch.randn(B, T, H, K, dtype=torch.float64, requires_grad=True)
k = torch.randn(B, T, H, K, dtype=torch.float64, requires_grad=True)
v = torch.randn(B, T, HV, V, dtype=torch.float64, requires_grad=True)
g = torch.randn(B, T, HV, K, dtype=torch.float64, requires_grad=True) * 0.1
beta = torch.rand(B, T, HV, dtype=torch.float64, requires_grad=True)
assert torch.autograd.gradcheck(
lambda q, k, v, g, b: naive_kda(q, k, v, g, b, output_final_state=True),
(q, k, v, g, beta),
eps=1e-6, atol=1e-4, rtol=1e-3,
), "L1 gradcheck 失败"
print("L1 gradcheck: PASSED")
if __name__ == "__main__":
test_gradcheck()
+1
View File
@@ -0,0 +1 @@
"""Incremental inference tests."""
+30
View File
@@ -0,0 +1,30 @@
"""L6: fused_recurrent decode matches naive recurrent."""
import pytest
import torch
from kda.ops.recurrent.fused import fused_recurrent_kda
from kda.ops.reference.recurrent import naive_kda
pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
def test_step_matches_naive():
B, T, H, HV, K, V = 2, 32, 2, 4, 8, 8
torch.manual_seed(20)
device = "cuda"
q = torch.randn(B, T, H, K, device=device)
k = torch.randn(B, T, H, K, device=device)
v = torch.randn(B, T, HV, V, device=device)
g = -torch.rand(B, T, HV, K, device=device) * 2
beta = torch.rand(B, T, HV, device=device)
o_naive, _ = naive_kda(
q.double(), k.double(), v.double(), g.double(),
beta.double(), output_final_state=False,
)
o_step, _ = fused_recurrent_kda(q, k, v, g, beta, output_final_state=False)
torch.testing.assert_close(o_step.float(), o_naive.float(), rtol=2e-3, atol=2e-3)
if __name__ == "__main__":
test_step_matches_naive()
+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()
+1
View File
@@ -0,0 +1 @@
"""Local kernel implementation tests."""
+20
View File
@@ -0,0 +1,20 @@
"""The vendored FLA fused gate must match the PyTorch reference."""
import pytest
import torch
from kda.ops.reference.gate import kda_gate_reference
from kda.ops.triton.gate import kda_gate_fwd
pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
def test_triton_gate_matches_reference():
B, T, HV, K = 1, 32, 4, 8
torch.manual_seed(10)
device = "cuda"
g = torch.randn(B, T, HV, K, device=device, dtype=torch.float32)
A_log = torch.randn(HV, device=device, dtype=torch.float32) * 0.5
dt_bias = torch.randn(HV, K, device=device, dtype=torch.float32) * 0.1
expected = kda_gate_reference(g, A_log, dt_bias)
actual = kda_gate_fwd(g, A_log, dt_bias, lower_bound=None)
torch.testing.assert_close(actual, expected, rtol=1e-4, atol=1e-4)
+67
View File
@@ -0,0 +1,67 @@
"""L4: vendored FLA Triton bwd vs naive chunked.
Triton kernels run in fp32, so this checks VJP vs L2 rather than fp64 gradcheck.
Inputs L2-normalize q/k like the trained KDA path.
"""
import pytest
import torch
import torch.nn.functional as F
from kda.ops.reference.chunkwise import naive_chunk_kda
from kda.ops.triton.chunk_fwd import chunk_kda_fwd
pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
def test_bwd_matches_naive():
B, T, H, HV, K, V = 2, 64, 4, 8, 32, 32
torch.manual_seed(5)
dev = "cuda"
q = F.normalize(torch.randn(B, T, H, K, device=dev), dim=-1).requires_grad_()
k = F.normalize(torch.randn(B, T, H, K, device=dev), dim=-1).requires_grad_()
v = torch.randn(B, T, HV, V, device=dev, requires_grad=True)
g = (-torch.rand(B, T, HV, K, device=dev) * 2).requires_grad_()
beta = torch.rand(B, T, HV, device=dev, requires_grad=True)
o_na, _ = naive_chunk_kda(q, k, v, g, beta, chunk_size=64)
o_na.sum().backward()
dq_na, dk_na = q.grad.clone(), k.grad.clone()
dv_na, dg_na, db_na = v.grad.clone(), g.grad.clone(), beta.grad.clone()
q.grad = k.grad = v.grad = g.grad = beta.grad = None
o_tr, _ = chunk_kda_fwd(q, k, v, g, beta, chunk_size=64)
o_tr.sum().backward()
dq_tr, dk_tr = q.grad.clone(), k.grad.clone()
dv_tr, dg_tr, db_tr = v.grad.clone(), g.grad.clone(), beta.grad.clone()
torch.testing.assert_close(dq_tr, dq_na, rtol=2e-2, atol=2e-3)
torch.testing.assert_close(dk_tr, dk_na, rtol=2e-2, atol=2e-3)
torch.testing.assert_close(dv_tr, dv_na, rtol=2e-2, atol=2e-3)
torch.testing.assert_close(dg_tr, dg_na, rtol=2e-2, atol=2e-3)
torch.testing.assert_close(db_tr, db_na, rtol=2e-2, atol=2e-3)
def test_triton_kda_attention_backward_dt_bias_rank2():
"""Layer stores dt_bias as [HV, K]; Triton bwd used to return a flat [HV*K]."""
from kda.layers.kda_attn import KDAAttention
torch.manual_seed(0)
model = KDAAttention(
hidden_size=64,
num_heads=4,
num_value_heads=8,
head_dim=16,
chunk_size=64,
kda_backend="triton",
).cuda()
x = torch.randn(2, 64, 64, device="cuda")
model(x).sum().backward()
assert model.dt_bias.grad is not None
assert model.dt_bias.grad.shape == model.dt_bias.shape
assert model.A_log.grad is not None
assert torch.isfinite(model.dt_bias.grad).all()
if __name__ == "__main__":
test_bwd_matches_naive()
test_triton_kda_attention_backward_dt_bias_rank2()
+29
View File
@@ -0,0 +1,29 @@
"""L3: vendored FLA Triton fwd vs naive chunked."""
import pytest
import torch
import torch.nn.functional as F
from kda.ops.reference.chunkwise import naive_chunk_kda
from kda.ops.triton.chunk_fwd import chunk_kda_fwd
pytestmark = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
def test_triton_fwd_matches_naive():
B, T, H, HV, K, V = 2, 64, 4, 8, 32, 32
torch.manual_seed(3)
device = "cuda"
q = F.normalize(torch.randn(B, T, H, K, device=device), dim=-1)
k = F.normalize(torch.randn(B, T, H, K, device=device), dim=-1)
v = torch.randn(B, T, HV, V, device=device)
g = -torch.rand(B, T, HV, K, device=device) * 2
beta = torch.rand(B, T, HV, device=device)
o_ref, S_ref = naive_chunk_kda(
q, k, v, g, beta, output_final_state=True, chunk_size=64
)
o_tr, S_tr = chunk_kda_fwd(
q, k, v, g, beta, output_final_state=True, chunk_size=64
)
torch.testing.assert_close(o_tr.float(), o_ref.float(), rtol=2e-3, atol=2e-3)
torch.testing.assert_close(S_tr.float(), S_ref.float(), rtol=2e-3, atol=2e-3)