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
+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
)