- replace the per-shape auto tables in the linear backend with an M-banded rule (M in [2,4] on compute capability 8.0+) that measured at the HBM bandwidth floor across every family, and fold the capability check into the capable guard - drop the unreachable swiglu auto shape-table machinery so both backends share one env-mode ladder via the new dispatch.env_mode helper - add __all__ across extension modules, name the rotary registration records, and unify typing to the typing-module style - rewrite test_linear_dispatch.py around behavioral routing assertions and document the M-banded policy in the developer docs - Benchmark: L20 SM89, Python dispatch overhead 2.9us to 1.5us, auto now covers every projection shape at M in [2,4].
35 lines
1000 B
Python
35 lines
1000 B
Python
"""Rotary embedding CUDA kernel wrapper.
|
|
|
|
Calls the compiled CUDA kernel directly. If the kernel is not available,
|
|
raises ``RuntimeError``. Fallback to torch complex multiply is the
|
|
responsibility of ``astrai.extension.backend.rotary.apply_rotary_emb``.
|
|
|
|
Layout: x is packed [tokens, n_heads, head_dim] or dense
|
|
[batch, seq_len, n_heads, head_dim]. ``freqs_cis`` has matching token axes.
|
|
"""
|
|
|
|
import torch
|
|
|
|
from astrai.extension.loader import get_module
|
|
|
|
|
|
def rotary_emb(x: torch.Tensor, freqs_cis: torch.Tensor) -> torch.Tensor:
|
|
"""Fused rotary embedding kernel.
|
|
|
|
Args:
|
|
x: packed 3D or dense 4D bf16 tensor.
|
|
freqs_cis: matching token axes followed by [head_dim/2, 2].
|
|
|
|
Returns:
|
|
Tensor with the same shape as ``x``.
|
|
"""
|
|
mod = get_module("rotary_emb")
|
|
if not x.is_contiguous():
|
|
x = x.contiguous()
|
|
if not freqs_cis.is_contiguous():
|
|
freqs_cis = freqs_cis.contiguous()
|
|
return mod.rotary_emb(x, freqs_cis)
|
|
|
|
|
|
__all__ = ["rotary_emb"]
|