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:
@@ -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
|
||||
)
|
||||
Reference in New Issue
Block a user