Sol-Attn-Kernel-Source / tests /test_sol_attn.py
hp-l33's picture
Add Sol-Attn Kernel Builder source
8e9f35a verified
Raw
History Blame Contribute Delete
2.64 kB
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)