"""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_{m1, 精确解本身就以 ‖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)