perf: add fused CUDA rotary embedding kernel
- 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%)
This commit is contained in:
@@ -0,0 +1,46 @@
|
||||
"""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.model.components.rope.apply_rotary_emb``.
|
||||
|
||||
Layout convention: x is ``[batch, seq_len, n_heads, head_dim]`` (blhd, bf16).
|
||||
cos/sin are ``[batch, seq_len, head_dim/2]`` (f32).
|
||||
"""
|
||||
|
||||
import torch
|
||||
|
||||
from astrai.extension.loader import _available, _modules
|
||||
|
||||
|
||||
def _check_available():
|
||||
if not _available.get("rotary_emb"):
|
||||
raise RuntimeError(
|
||||
"CUDA kernel 'rotary_emb' is not available. "
|
||||
"Build with CSRC_KERNELS=true or use the torch fallback."
|
||||
)
|
||||
|
||||
|
||||
def rotary_emb(
|
||||
x: torch.Tensor,
|
||||
cos: torch.Tensor,
|
||||
sin: torch.Tensor,
|
||||
) -> torch.Tensor:
|
||||
"""Fused rotary embedding kernel.
|
||||
|
||||
Applies rotation: for each pair (x_even, x_odd):
|
||||
out_even = x_even * cos - x_odd * sin
|
||||
out_odd = x_even * sin + x_odd * cos
|
||||
|
||||
Args:
|
||||
x: [batch, seq_len, n_heads, head_dim] (bf16, contiguous)
|
||||
cos: [batch, seq_len, head_dim/2] (f32)
|
||||
sin: [batch, seq_len, head_dim/2] (f32)
|
||||
|
||||
Returns:
|
||||
[batch, seq_len, n_heads, head_dim] (bf16)
|
||||
"""
|
||||
_check_available()
|
||||
if not x.is_contiguous():
|
||||
x = x.contiguous()
|
||||
return _modules["rotary_emb"].rotary_emb(x, cos, sin)
|
||||
Reference in New Issue
Block a user