- Single-kernel rotary embedding (cos/sin lookup + rotation) replaces PyTorch complex-multiply path (3 kernel launches + f32 upcast per call) - RotaryEmbedding now stores cos_table/sin_table and returns (cos, sin) f32 tuple instead of a complex tensor - apply_rotary_emb in rotary_backend.py auto-dispatches: CUDA kernel if available, else torch complex-multiply fallback; backend-agnostic (both attention backends benefit) - Kernel: 256-thread blocks, grid-stride loop, vectorized __nv_bfloat162 load/store, f32 compute, bf16 out - Standalone kernel 6-9x faster than torch across decode/prefill shapes, max diff 0 (decode) to 3e-2 (large prefill, bf16) - Benchmark (L20, bf16, CUDA backend): B=1 9.48->7.25ms (+31%), B=4 10.73->7.67ms (+40%), B=8 10.77->7.81ms (+38%), B=16 10.79->7.83ms (+38%)
37 lines
1.1 KiB
Python
37 lines
1.1 KiB
Python
"""Dynamic discovery and loading of compiled CUDA kernel modules.
|
|
|
|
Each kernel is registered in ``csrc/build.py`` and built into a ``.so`` placed
|
|
in this package directory. On import we try to load each one; kernels that
|
|
failed to build (or are running on a CPU-only machine) are marked unavailable
|
|
so the wrapper functions can fall back to ``torch`` SDPA.
|
|
"""
|
|
|
|
import importlib
|
|
import logging
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
KERNEL_NAMES = ["attn_decode", "attn_prefill", "attn_paged_decode", "rotary_emb"]
|
|
|
|
_available: dict[str, bool] = {}
|
|
_modules: dict[str, object] = {}
|
|
|
|
for _name in KERNEL_NAMES:
|
|
try:
|
|
_mod = importlib.import_module(f".{_name}", package=__package__)
|
|
_available[_name] = True
|
|
_modules[_name] = _mod
|
|
except ImportError:
|
|
_available[_name] = False
|
|
_modules[_name] = None
|
|
|
|
|
|
def is_available(name: str) -> bool:
|
|
"""Return ``True`` if the compiled kernel ``name`` was loaded."""
|
|
return _available.get(name, False)
|
|
|
|
|
|
def get_module(name: str) -> object:
|
|
"""Return the loaded kernel module for ``name``, or ``None`` if unavailable."""
|
|
return _modules.get(name)
|