diff --git a/astrai/extension/__init__.py b/astrai/extension/__init__.py index 8ec0bab..b39aaaf 100644 --- a/astrai/extension/__init__.py +++ b/astrai/extension/__init__.py @@ -30,6 +30,7 @@ from astrai.extension.attention_ops import ( attn_prefill, ) from astrai.extension.loader import KERNEL_NAMES, is_available +from astrai.extension.rotary_backend import apply_rotary_emb __all__ = [ "ATTN_BACKEND", @@ -44,4 +45,5 @@ __all__ = [ "attn_prefill", "is_available", "KERNEL_NAMES", + "apply_rotary_emb", ] diff --git a/astrai/extension/loader.py b/astrai/extension/loader.py index 63764f4..587ac3c 100644 --- a/astrai/extension/loader.py +++ b/astrai/extension/loader.py @@ -11,7 +11,7 @@ import logging logger = logging.getLogger(__name__) -KERNEL_NAMES = ["attn_decode", "attn_prefill", "attn_paged_decode"] +KERNEL_NAMES = ["attn_decode", "attn_prefill", "attn_paged_decode", "rotary_emb"] _available: dict[str, bool] = {} _modules: dict[str, object] = {} diff --git a/astrai/extension/rotary_backend.py b/astrai/extension/rotary_backend.py new file mode 100644 index 0000000..3803a32 --- /dev/null +++ b/astrai/extension/rotary_backend.py @@ -0,0 +1,48 @@ +"""Rotary embedding with auto-dispatch to CUDA kernel. + +Single entry point ``apply_rotary_emb(x, cos, sin)`` — uses the fused +CUDA kernel when available, falls back to torch complex multiply otherwise. + +Layout: x is [batch, seq_len, n_heads, head_dim] (bf16). +cos/sin are [batch, seq_len, head_dim/2] (f32). +""" + +import torch +from torch import Tensor + +from astrai.extension.loader import is_available +from astrai.extension.rotary_ops import rotary_emb as _cuda_rotary + +_cache = {"available": None} + + +def _cuda_available() -> bool: + if _cache["available"] is None: + _cache["available"] = is_available("rotary_emb") + return _cache["available"] + + +def _torch_apply(x: Tensor, cos: Tensor, sin: Tensor) -> Tensor: + dtype = x.dtype + x_ = x.float().reshape(*x.shape[:-1], -1, 2) + x_complex = torch.view_as_complex(x_) + freqs_cis = torch.complex(cos, sin).unsqueeze(2) + x_rotated = x_complex * freqs_cis + x_out = torch.view_as_real(x_rotated).flatten(-2) + return x_out.to(dtype) + + +def apply_rotary_emb(x: Tensor, rotary_emb: tuple[Tensor, Tensor]) -> Tensor: + """Apply rotary embedding to x. + + Args: + x: [batch, seq_len, n_heads, head_dim] (bf16) + rotary_emb: (cos, sin) tuple, each [batch, seq_len, head_dim/2] (f32) + + Returns: + [batch, seq_len, n_heads, head_dim] (bf16) + """ + cos, sin = rotary_emb + if _cuda_available() and x.is_cuda and x.dtype == torch.bfloat16: + return _cuda_rotary(x, cos, sin) + return _torch_apply(x, cos, sin) diff --git a/astrai/extension/rotary_ops.py b/astrai/extension/rotary_ops.py new file mode 100644 index 0000000..512d894 --- /dev/null +++ b/astrai/extension/rotary_ops.py @@ -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) diff --git a/astrai/model/components/__init__.py b/astrai/model/components/__init__.py index f304af2..96e0cda 100644 --- a/astrai/model/components/__init__.py +++ b/astrai/model/components/__init__.py @@ -1,3 +1,4 @@ +from astrai.extension.rotary_backend import apply_rotary_emb from astrai.model.components.attention import GQA, MLA from astrai.model.components.decoder_block import DecoderBlock from astrai.model.components.embedding import Embedding @@ -6,7 +7,6 @@ from astrai.model.components.mlp import MLP from astrai.model.components.norm import RMSNorm from astrai.model.components.rope import ( RotaryEmbedding, - apply_rotary_emb, get_rotary_emb, ) diff --git a/astrai/model/components/attention.py b/astrai/model/components/attention.py index 8cf217b..4428ad0 100644 --- a/astrai/model/components/attention.py +++ b/astrai/model/components/attention.py @@ -6,11 +6,11 @@ import torch.nn.functional as F from torch import Tensor from astrai.extension import attention +from astrai.extension.rotary_backend import apply_rotary_emb from astrai.factory import BaseFactory from astrai.inference.core.cache import KVCache from astrai.model.components.linear import Linear from astrai.model.components.norm import RMSNorm -from astrai.model.components.rope import apply_rotary_emb class AttnFactory(BaseFactory[nn.Module]): diff --git a/astrai/model/components/rope.py b/astrai/model/components/rope.py index bbffb61..2f4264e 100644 --- a/astrai/model/components/rope.py +++ b/astrai/model/components/rope.py @@ -1,4 +1,4 @@ -from typing import Dict, Optional +from typing import Dict, Optional, Tuple import torch import torch.nn as nn @@ -10,29 +10,22 @@ def get_rotary_emb( max_len: int, base: float = 10000, device: Optional[torch.device] = None, -) -> Tensor: +) -> Tuple[Tensor, Tensor]: + """Precompute cos/sin tables for rotary embedding. + + Returns: + (cos, sin) each of shape [max_len, dim/2] (f32) + """ theta = base ** (-torch.arange(0, dim, 2, dtype=torch.float64, device=device) / dim) t = torch.arange(0, max_len, dtype=torch.float64, device=device) freqs = torch.outer(t, theta).float() - cos = torch.cos(freqs) - sin = torch.sin(freqs) - return torch.complex(cos, sin) + return torch.cos(freqs), torch.sin(freqs) def ntk_base(base: float, dim: int, factor: float) -> float: return base * (factor ** (dim / (dim - 2))) -def apply_rotary_emb(x: torch.Tensor, freqs_cis: Tensor) -> Tensor: - dtype = x.dtype - x_ = x.float().reshape(*x.shape[:-1], -1, 2) - x_complex = torch.view_as_complex(x_) - freqs_cis = freqs_cis.unsqueeze(2) - x_rotated = x_complex * freqs_cis - x_out = torch.view_as_real(x_rotated).flatten(-2) - return x_out.to(dtype) - - class RotaryEmbedding(nn.Module): def __init__( self, @@ -56,16 +49,28 @@ class RotaryEmbedding(nn.Module): self._set_rotary_buffer(self.max_len) def _set_rotary_buffer(self, max_len: int): - rotary_emb = get_rotary_emb(self.dim, max_len, self.base) - freqs_cis = torch.view_as_real(rotary_emb) - self.register_buffer("freqs_cis", freqs_cis, persistent=False) + cos, sin = get_rotary_emb(self.dim, max_len, self.base) + self.register_buffer("cos_table", cos, persistent=False) + self.register_buffer("sin_table", sin, persistent=False) - def forward(self, x: Tensor, position_ids: Optional[Tensor] = None) -> Tensor: + def forward( + self, x: Tensor, position_ids: Optional[Tensor] = None + ) -> Tuple[Tensor, Tensor]: + """Lookup cos/sin for the given positions. + + Args: + x: [batch, seq_len, ...] — only batch and seq_len are used. + position_ids: [batch, seq_len] optional position indices. + + Returns: + (cos, sin) each of shape [batch, seq_len, dim/2] (f32) + """ if position_ids is None: position_ids = ( torch.arange(x.size(1), device=x.device) .unsqueeze(0) .expand(x.size(0), -1) ) - position_freq_cis = self.freqs_cis[position_ids].float() - return torch.view_as_complex(position_freq_cis) + cos = self.cos_table[position_ids].float() + sin = self.sin_table[position_ids].float() + return cos, sin diff --git a/csrc/build.py b/csrc/build.py index 7f1d1ab..a5f78c7 100644 --- a/csrc/build.py +++ b/csrc/build.py @@ -72,3 +72,4 @@ def register(name: str, sources: list[str] | None = None, **kwargs): register("attn_decode") register("attn_prefill") register("attn_paged_decode") +register("rotary_emb") diff --git a/csrc/kernels/rotary_emb.cu b/csrc/kernels/rotary_emb.cu new file mode 100644 index 0000000..ef763c8 --- /dev/null +++ b/csrc/kernels/rotary_emb.cu @@ -0,0 +1,92 @@ +#include +#include + +__global__ void rotary_emb_kernel( + const __nv_bfloat16* __restrict__ x, + const float* __restrict__ cos, + const float* __restrict__ sin, + __nv_bfloat16* __restrict__ out, + int batch, + int seq_len, + int n_heads, + int head_dim +) { + const int half_dim = head_dim >> 1; + const int total = batch * seq_len * n_heads * half_dim; + + for (int idx = blockIdx.x * blockDim.x + threadIdx.x; + idx < total; + idx += gridDim.x * blockDim.x) { + + int pair = idx % half_dim; + int tmp = idx / half_dim; + int head = tmp % n_heads; + tmp /= n_heads; + int seq = tmp % seq_len; + int b = tmp / seq_len; + + int x_offset = ((b * seq_len + seq) * n_heads + head) * head_dim + (pair << 1); + int cs_offset = (b * seq_len + seq) * half_dim + pair; + + __nv_bfloat162 x_pair = *reinterpret_cast(x + x_offset); + float x_even = __bfloat162float(__low2bfloat16(x_pair)); + float x_odd = __bfloat162float(__high2bfloat16(x_pair)); + + float c = cos[cs_offset]; + float s = sin[cs_offset]; + + float out_even = x_even * c - x_odd * s; + float out_odd = x_even * s + x_odd * c; + + __nv_bfloat162 out_pair = __floats2bfloat162_rn(out_even, out_odd); + *reinterpret_cast<__nv_bfloat162*>(out + x_offset) = out_pair; + } +} + +torch::Tensor rotary_emb( + torch::Tensor x, + torch::Tensor cos, + torch::Tensor sin +) { + + int batch = x.size(0); + int seq_len = x.size(1); + int n_heads = x.size(2); + int head_dim = x.size(3); + + TORCH_CHECK(x.is_cuda(), "x must be on CUDA"); + TORCH_CHECK(cos.is_cuda(), "cos must be on CUDA"); + TORCH_CHECK(sin.is_cuda(), "sin must be on CUDA"); + TORCH_CHECK(x.scalar_type() == torch::kBFloat16, "x must be bf16"); + TORCH_CHECK(x.dim() == 4, "x must be 4D [batch, seq_len, n_heads, head_dim]"); + TORCH_CHECK(x.is_contiguous(), "x must be contiguous"); + TORCH_CHECK(cos.dim() == 3, "cos must be 3D [batch, seq_len, head_dim/2]"); + TORCH_CHECK(sin.dim() == 3, "sin must be 3D [batch, seq_len, head_dim/2]"); + TORCH_CHECK(head_dim % 2 == 0, "head_dim must be even"); + + auto out = torch::empty_like(x); + + int half_dim = head_dim / 2; + int total = batch * seq_len * n_heads * half_dim; + int block = 256; + int grid = std::min((total + block - 1) / block, 1024); + + rotary_emb_kernel<<>>( + reinterpret_cast(x.data_ptr()), + cos.data_ptr(), + sin.data_ptr(), + reinterpret_cast<__nv_bfloat16*>(out.data_ptr()), + batch, seq_len, n_heads, head_dim + ); + + return out; +} + +PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { + m.def("rotary_emb", &rotary_emb, + py::arg("x"), + py::arg("cos"), + py::arg("sin"), + "Fused rotary embedding (bf16 x, f32 cos/sin, bf16 out)" + ); +}