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