Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
100 lines
3.6 KiB
Python
100 lines
3.6 KiB
Python
"""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
|
|
)
|