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 @@
"""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()