refactor: separate extension ops and backends

This commit is contained in:
2026-08-16 21:15:52 +08:00
parent 3406157431
commit 6ac3b51496
18 changed files with 82 additions and 54 deletions
+19
View File
@@ -0,0 +1,19 @@
"""Stateless wrappers around compiled extension kernels."""
from astrai.extension.ops.attention import (
TensorLayout,
attn_decode,
attn_paged_decode,
attn_paged_prefill,
attn_prefill,
)
from astrai.extension.ops.rotary import rotary_emb
__all__ = [
"TensorLayout",
"attn_decode",
"attn_paged_decode",
"attn_paged_prefill",
"attn_prefill",
"rotary_emb",
]
+188
View File
@@ -0,0 +1,188 @@
"""Attention kernel wrapper functions - one entry point per compiled kernel.
Each wrapper calls its CUDA kernel directly. If the kernel is not
available, raises ``RuntimeError``. Fallback to torch SDPA is the
responsibility of the attention backend, not this module.
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
(blhd). Scale is always ``1/sqrt(head_dim)``.
Interface (all functions):
is_causal: True = causal mask; False = non-causal
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
"""
import enum
from typing import Optional
import torch
from astrai.extension.loader import _available, _modules
class TensorLayout(enum.IntEnum):
"""Q/K/V tensor layout, mirrors the C++ ``TensorLayout`` enum in ``attn_common.h``.
Kernels internally operate on BHLD; BLHD inputs are transposed at entry.
"""
BHLD = 0 # [batch, n_heads, seq_len, head_dim]
BLHD = 1 # [batch, seq_len, n_heads, head_dim]
def _check_available(name: str):
if not _available.get(name):
raise RuntimeError(
f"CUDA kernel '{name}' is not available. "
f"Build with CSRC_KERNELS=true or use a torch-native backend."
)
def attn_decode(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: Optional[torch.Tensor] = None,
is_causal: bool = False,
) -> torch.Tensor:
"""GQA decode attention (q_len == 1).
Args:
q: [batch, 1, n_heads, head_dim] (blhd, bf16)
k: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
v: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
mask: 2D [batch, kv_len] or 3D [batch, 1, kv_len] (bool, True=keep)
is_causal: apply causal mask
Returns:
[batch, 1, n_heads, head_dim] (blhd, bf16)
"""
_check_available("attn_decode")
causal_offset = (k.size(1) - 1) if is_causal else -1
return _modules["attn_decode"].attn_decode(
q, k, v, mask=mask, causal_offset=causal_offset, layout=TensorLayout.BLHD
)
def attn_prefill(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: Optional[torch.Tensor] = None,
is_causal: bool = False,
) -> torch.Tensor:
"""GQA prefill attention (q_len > 1).
Args:
q: [batch, q_len, n_heads, head_dim] (blhd, bf16)
k: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
v: [batch, kv_len, n_kv_heads, head_dim] (blhd, bf16)
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True=keep)
is_causal: apply causal mask
Returns:
[batch, q_len, n_heads, head_dim] (blhd, bf16)
"""
_check_available("attn_prefill")
causal_offset = (k.size(1) - q.size(1)) if is_causal else -1
return _modules["attn_prefill"].attn_prefill(
q, k, v, mask=mask, causal_offset=causal_offset, layout=TensorLayout.BLHD
)
def attn_paged_decode(
q: torch.Tensor,
k_cache: torch.Tensor,
v_cache: torch.Tensor,
req_to_token: torch.Tensor,
req_pool_indices: torch.Tensor,
kv_indptr: torch.Tensor,
mask: Optional[torch.Tensor] = None,
is_causal: bool = False,
o_part_buf: Optional[torch.Tensor] = None,
ml_part_buf: Optional[torch.Tensor] = None,
out_buf: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""SGLang-style paged decode (q_len == 1, flat KV pool).
Reads K/V directly from a flat pool [size, kv_head, head_dim] via
req_to_token indirect indexing. Each request has its own seq_len
(from kv_indptr), eliminating padding waste.
Args:
q: [batch, n_heads, head_dim] (bf16, 3D — no seq dim)
k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat)
v_cache: same as k_cache
req_to_token: [num_reqs, max_context_len] (int32) — token -> slot
req_pool_indices: [batch] (int32) — rows into req_to_token
kv_indptr: [batch+1] (int32) — prefix sum of per-request seq_lens
mask: 2D [batch, max_context_len] (bool, True=keep) or None
is_causal: apply causal mask
o_part_buf: pre-allocated split-KV o partial buffer (workflow bypass)
ml_part_buf: pre-allocated split-KV m/l buffer (workflow bypass)
out_buf: pre-allocated output buffer [batch, n_heads, head_dim] (graph-safe)
Returns:
[batch, n_heads, head_dim] (bf16, 3D)
"""
_check_available("attn_paged_decode")
causal_offset = 0 if is_causal else -1
return _modules["attn_paged_decode"].attn_paged_decode(
q,
k_cache,
v_cache,
req_to_token,
req_pool_indices,
kv_indptr,
mask=mask,
causal_offset=causal_offset,
o_part_buf=o_part_buf,
ml_part_buf=ml_part_buf,
out_buf=out_buf,
)
def attn_paged_prefill(
q: torch.Tensor,
k_cache: torch.Tensor,
v_cache: torch.Tensor,
req_to_token: torch.Tensor,
req_pool_indices: torch.Tensor,
kv_indptr: torch.Tensor,
qo_indptr: torch.Tensor,
mask: Optional[torch.Tensor] = None,
is_causal: bool = False,
) -> torch.Tensor:
"""SGLang-style paged prefill (ragged batch, flat KV pool).
Reads K/V directly from a flat pool [size, kv_head, head_dim] via
req_to_token. Supports ragged batches: each request has its own
q_len and kv_len, addressed via qo_indptr and kv_indptr.
Args:
q: [total_q, n_heads, head_dim] (bf16, 3D — flattened across requests)
k_cache: [pool_size, n_kv_heads, head_dim] (bf16, flat)
v_cache: same as k_cache
req_to_token: [num_reqs, max_context_len] (int32)
req_pool_indices: [batch] (int32)
kv_indptr: [batch+1] (int32) — prefix sum of per-request kv_lens
qo_indptr: [batch+1] (int32) — prefix sum of per-request q_lens
mask: 4D [batch, 1, q_len, kv_len] (bool, True=keep) or None
is_causal: apply causal mask
Returns:
[total_q, n_heads, head_dim] (bf16, 3D)
"""
_check_available("attn_paged_prefill")
causal_offset = 0 if is_causal else -1
return _modules["attn_paged_prefill"].attn_paged_prefill(
q,
k_cache,
v_cache,
req_to_token,
req_pool_indices,
kv_indptr,
qo_indptr,
mask,
causal_offset=causal_offset,
)
+73
View File
@@ -0,0 +1,73 @@
"""FP8 CUDA kernel interface adapter (the only module touching the pybind.
Isolates the ``fp8_mm`` CUDA extension behind stable Python functions:
- availability / dtype checks and clear errors
- torch.library ``custom::fp8_mm`` registration (meta + CPU fallback)
- quantize-in-GEMM primitives used by ``fp8.py`` training state
Policy (scales, amax history, delayed scaling, autocast) lives in ``fp8.py``;
this module is stateless.
"""
import torch
from torch.library import custom_op
from astrai.extension.loader import get_module, is_available
def _mod():
if not is_available("fp8_mm"):
raise RuntimeError(
"CUDA kernel 'fp8_mm' is not available. Build with CSRC_KERNELS=true."
)
return get_module("fp8_mm")
@custom_op("custom::fp8_mm", mutates_args=())
def fp8_mm(
a: torch.Tensor, b: torch.Tensor, sx: torch.Tensor, sw: torch.Tensor
) -> torch.Tensor:
"""FP8 e4m3 GEMM: a[M,K] x b[N,K] -> bf16[M,N] (pre-scaled inputs)."""
@fp8_mm.register_fake
def _fp8_mm_fake(a, b, sx, sw):
return torch.empty((a.size(0), b.size(1)), device=a.device, dtype=torch.bfloat16)
@fp8_mm.register_kernel("cuda")
def _fp8_mm_cuda(a, b, sx, sw):
return _mod().fp8_mm(a, b)
@fp8_mm.register_kernel("cpu")
def _fp8_mm_cpu(a, b, sx, sw):
return torch.mm(a.float(), b.float().t()).to(torch.bfloat16)
def linear_forward_scaled(x, w, bias, sx, sw, sx_inv, sw_inv, amax_x, amax_w):
"""Quantize x/w with per-tensor scales + cuBLASLt GEMM + bias -> bf16.
x/w: [..., K] / [N, K] bf16; sx/sw: f32 scale tensors (device scalars);
sx_inv/sw_inv: 1/scale; amax_x/amax_w: f32 buffers receiving max-abs.
"""
if not (x.dtype == torch.bfloat16 and w.dtype == torch.bfloat16):
raise TypeError(f"fp8 forward requires bf16 inputs, got {x.dtype}/{w.dtype}")
return _mod().fp8_linear_forward_scaled(
x, w, bias, sx, sw, sx_inv, sw_inv, amax_x, amax_w
)
def linear_backward_scaled(g, x, w, masks, sg, sw, sx, sg_inv, sw_inv, sx_inv, amax_g):
"""dX = g @ W, dW = g^T @ X, dB = sum(g) with per-tensor scales."""
if not (
g.dtype == torch.bfloat16
and x.dtype == torch.bfloat16
and w.dtype == torch.bfloat16
):
raise TypeError(
f"fp8 backward requires bf16 inputs, got {g.dtype}/{x.dtype}/{w.dtype}"
)
return _mod().fp8_linear_backward_scaled(
g, x, w, masks, sg, sw, sx, sg_inv, sw_inv, sx_inv, amax_g
)
+39
View File
@@ -0,0 +1,39 @@
"""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 _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, 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``.
"""
_check_available()
if not x.is_contiguous():
x = x.contiguous()
if not freqs_cis.is_contiguous():
freqs_cis = freqs_cis.contiguous()
return _modules["rotary_emb"].rotary_emb(x, freqs_cis)