import importlib import pytest import torch import torch.nn.functional as F from kernels import get_kernel kernel = get_kernel("Efficient-Large-Model/Sol-Attn", version=1) def _inputs(tokens=256, heads=4): torch.manual_seed(42) q = torch.randn( 1, tokens, heads, 128, device="cuda", dtype=torch.bfloat16, ) return q, torch.randn_like(q), torch.randn_like(q) @pytest.mark.kernels_ci def test_backend_dispatch_contract(): interface = importlib.import_module(f"{kernel.__name__}.interface") assert interface._backend_for_arch((8, 0), cute_available=True) == "triton" assert interface._backend_for_arch((8, 9), cute_available=True) == "triton" assert interface._backend_for_arch((9, 0), cute_available=True) == "cute_sm90" assert interface._backend_for_arch((10, 0), cute_available=True) == "cute_sm100" assert interface._backend_for_arch((12, 0), cute_available=True) == "cute_sm120" assert interface._backend_for_arch((9, 0), cute_available=False) == "triton" assert interface._backend_for_arch((10, 0), cute_available=False) == "triton" assert interface._backend_for_arch((12, 0), cute_available=False) == "triton" @pytest.mark.kernels_ci @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") def test_full_sink_matches_sdpa(): capability = torch.cuda.get_device_capability() if capability[0] < 8: pytest.skip("Sol-Attn requires compute capability 8.0 or newer") q, k, v = _inputs() expected = F.scaled_dot_product_attention( q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), ).transpose(1, 2) actual = kernel.sol_attn( q, k, v, tau=1.0, thresh_type="exact", sink_start=0, sink_tokens=q.shape[1], ) torch.testing.assert_close(actual, expected, atol=2e-2, rtol=3e-2) @pytest.mark.kernels_ci @pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA is required") def test_selected_backend_matches_triton_reference(): capability = torch.cuda.get_device_capability() if capability[0] < 8: pytest.skip("Sol-Attn requires compute capability 8.0 or newer") triton_ref = importlib.import_module(f"{kernel.__name__}.triton_ref") q, k, v = _inputs() expected = triton_ref.sol_attn( q, k, v, tau=1.0, thresh_type="exact", ) actual = kernel.sol_attn( q, k, v, tau=1.0, thresh_type="exact", ) torch.testing.assert_close(actual, expected, atol=2e-2, rtol=3e-2)