Instructions to use Efficient-Large-Model/Sol-Attn-Kernel-Source with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Kernels
How to use Efficient-Large-Model/Sol-Attn-Kernel-Source with Kernels:
# !pip install kernels from kernels import get_kernel kernel = get_kernel("Efficient-Large-Model/Sol-Attn-Kernel-Source") - Notebooks
- Google Colab
- Kaggle
| 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) | |
| 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" | |
| 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) | |
| 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) | |