refactor: unify extension API to blhd layout and is_causal

- Rename ops.py to attention_ops.py
- Remove layout/scale params: fixed blhd, auto scale
- Replace causal_offset with is_causal bool
- Move SDPA fallback to backend, ops only calls CUDA kernels
- Update __init__.py exports
This commit is contained in:
2026-07-30 18:39:20 +08:00
parent 5b67d5865a
commit 21bf37dd83
4 changed files with 147 additions and 341 deletions
+21 -9
View File
@@ -4,27 +4,39 @@ Public API:
- ``attn_decode`` — single-query decode attention - ``attn_decode`` — single-query decode attention
- ``attn_prefill`` — multi-query prefill attention - ``attn_prefill`` — multi-query prefill attention
- ``attn_paged_decode`` — paged decode attention (direct page-table access) - ``attn_paged_decode`` — paged decode attention (direct page-table access)
- ``AttentionBackend`` — ABC for attention computation strategies
- ``TorchNativeBackend`` — default SDPA backend with KV cache I/O
Interface (shared by all wrappers): Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
causal_offset: -1 = non-causal; >=0 = absolute position of first Q token (blhd). Scale is always ``1/sqrt(head_dim)``.
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool, True = keep)
scale: 0.0 = auto (1/sqrt(head_dim)); >0 = explicit
layout: "bhld" (default) or "blhd"
Causal and mask can coexist — both are applied simultaneously.
Each wrapper dispatches to its compiled CUDA kernel (``astrai.extension.attn_*``) Each wrapper dispatches to its compiled CUDA kernel (``astrai.extension.attn_*``)
when available, otherwise falls back to ``torch.nn.functional.scaled_dot_product_attention``. when available, otherwise falls back to ``torch.nn.functional.scaled_dot_product_attention``.
""" """
from astrai.extension.attention_backend import (
ATTN_BACKEND,
AttentionBackend,
TorchNativeBackend,
attn_backend,
get_backend,
)
from astrai.extension.attention_ops import (
attn_decode,
attn_paged_decode,
attn_prefill,
)
from astrai.extension.loader import KERNEL_NAMES, is_available from astrai.extension.loader import KERNEL_NAMES, is_available
from astrai.extension.ops import attention, attn_decode, attn_paged_decode, attn_prefill
__all__ = [ __all__ = [
"ATTN_BACKEND",
"AttentionBackend",
"TorchNativeBackend",
"attn_backend",
"get_backend",
"attn_decode", "attn_decode",
"attn_paged_decode", "attn_paged_decode",
"attn_prefill", "attn_prefill",
"attention",
"is_available", "is_available",
"KERNEL_NAMES", "KERNEL_NAMES",
] ]
+9 -34
View File
@@ -24,15 +24,13 @@ Thread-safe via ``contextvars`` — each scheduler thread gets its own
active backend. ``get_backend()`` returns the active one, falling back active backend. ``get_backend()`` returns the active one, falling back
to a process-wide ``TorchNativeBackend`` singleton. to a process-wide ``TorchNativeBackend`` singleton.
Contract (q/k/v are always ``[batch, seq_len, n_heads, head_dim]``): Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
q: [batch, q_len, n_heads, head_dim] (blhd). The backend returns ``[batch, seq_len, n_heads * head_dim]``.
k: [batch, q_len, n_kv_heads, head_dim]
v: [batch, q_len, n_kv_heads, head_dim]
-> returns [batch, q_len, n_heads * head_dim]
""" """
import contextvars import contextvars
import enum import enum
import math
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
from contextlib import contextmanager from contextlib import contextmanager
from typing import Optional, Union from typing import Optional, Union
@@ -143,10 +141,8 @@ class AttentionBackend(ABC):
v: Tensor, v: Tensor,
kv_cache: Optional[KVCache], kv_cache: Optional[KVCache],
layer_id: int, layer_id: int,
n_rep: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
scale: Optional[float] = None,
) -> Tensor: ) -> Tensor:
"""Dispatch to decode or extend based on q_len. """Dispatch to decode or extend based on q_len.
@@ -156,21 +152,17 @@ class AttentionBackend(ABC):
v: [batch, q_len, n_kv_heads, head_dim] v: [batch, q_len, n_kv_heads, head_dim]
kv_cache: cache dataclass, or None for training (no cache). kv_cache: cache dataclass, or None for training (no cache).
layer_id: transformer layer index for buffer access. layer_id: transformer layer index for buffer access.
n_rep: Q-heads // KV-heads (1 when heads match, e.g. MLA).
attn_mask: pre-built attention mask compatible with SDPA. attn_mask: pre-built attention mask compatible with SDPA.
is_causal: whether to apply causal masking. is_causal: whether to apply causal masking.
scale: explicit softmax scale; None = 1/sqrt(head_dim).
Returns: Returns:
[batch, q_len, n_heads * head_dim] [batch, q_len, n_heads * head_dim]
""" """
if kv_cache is not None and q.size(1) == 1: if kv_cache is not None and q.size(1) == 1:
return self.forward_decode( return self.forward_decode(
q, k, v, kv_cache, layer_id, n_rep, attn_mask, is_causal, scale q, k, v, kv_cache, layer_id, attn_mask, is_causal
) )
return self.forward_extend( return self.forward_extend(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
q, k, v, kv_cache, layer_id, n_rep, attn_mask, is_causal, scale
)
@abstractmethod @abstractmethod
def forward_decode( def forward_decode(
@@ -180,10 +172,8 @@ class AttentionBackend(ABC):
v: Tensor, v: Tensor,
kv_cache: Optional[KVCache], kv_cache: Optional[KVCache],
layer_id: int, layer_id: int,
n_rep: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
scale: Optional[float] = None,
) -> Tensor: ) -> Tensor:
"""Single-token decode with KV cache.""" """Single-token decode with KV cache."""
@@ -195,10 +185,8 @@ class AttentionBackend(ABC):
v: Tensor, v: Tensor,
kv_cache: Optional[KVCache], kv_cache: Optional[KVCache],
layer_id: int, layer_id: int,
n_rep: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
scale: Optional[float] = None,
) -> Tensor: ) -> Tensor:
"""Multi-token prefill or training forward.""" """Multi-token prefill or training forward."""
@@ -221,14 +209,10 @@ class TorchNativeBackend(AttentionBackend):
v: Tensor, v: Tensor,
kv_cache: Optional[KVCache], kv_cache: Optional[KVCache],
layer_id: int, layer_id: int,
n_rep: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
scale: Optional[float] = None,
) -> Tensor: ) -> Tensor:
return self._forward( return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
q, k, v, kv_cache, layer_id, n_rep, attn_mask, is_causal, scale
)
def forward_extend( def forward_extend(
self, self,
@@ -237,14 +221,10 @@ class TorchNativeBackend(AttentionBackend):
v: Tensor, v: Tensor,
kv_cache: Optional[KVCache], kv_cache: Optional[KVCache],
layer_id: int, layer_id: int,
n_rep: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
scale: Optional[float] = None,
) -> Tensor: ) -> Tensor:
return self._forward( return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
q, k, v, kv_cache, layer_id, n_rep, attn_mask, is_causal, scale
)
def _forward( def _forward(
self, self,
@@ -253,10 +233,8 @@ class TorchNativeBackend(AttentionBackend):
v: Tensor, v: Tensor,
kv_cache: Optional[KVCache], kv_cache: Optional[KVCache],
layer_id: int, layer_id: int,
n_rep: int,
attn_mask: Optional[Tensor] = None, attn_mask: Optional[Tensor] = None,
is_causal: bool = False, is_causal: bool = False,
scale: Optional[float] = None,
) -> Tensor: ) -> Tensor:
if kv_cache is not None: if kv_cache is not None:
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
@@ -272,6 +250,7 @@ class TorchNativeBackend(AttentionBackend):
k = kv_cache.k_buffer[layer_id, indices] k = kv_cache.k_buffer[layer_id, indices]
v = kv_cache.v_buffer[layer_id, indices] v = kv_cache.v_buffer[layer_id, indices]
n_rep = q.size(2) // k.size(2)
if n_rep > 1: if n_rep > 1:
k = repeat_kv(k, n_rep) k = repeat_kv(k, n_rep)
v = repeat_kv(v, n_rep) v = repeat_kv(v, n_rep)
@@ -280,11 +259,7 @@ class TorchNativeBackend(AttentionBackend):
k = k.permute(0, 2, 1, 3) k = k.permute(0, 2, 1, 3)
v = v.permute(0, 2, 1, 3) v = v.permute(0, 2, 1, 3)
sdpa_kwargs: dict = {"is_causal": is_causal} out = F.scaled_dot_product_attention(q, k, v, attn_mask, is_causal=is_causal)
if scale is not None:
sdpa_kwargs["scale"] = scale
out = F.scaled_dot_product_attention(q, k, v, attn_mask, **sdpa_kwargs)
out = out.permute(0, 2, 1, 3).contiguous().flatten(2) out = out.permute(0, 2, 1, 3).contiguous().flatten(2)
return out return out
+117
View File
@@ -0,0 +1,117 @@
"""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 torch
from astrai.extension.loader import _available, _modules
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: torch.Tensor | None = 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=1
)
def attn_prefill(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: torch.Tensor | None = 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=1
)
def attn_paged_decode(
q: torch.Tensor,
page_table: torch.Tensor,
k_cache: torch.Tensor,
v_cache: torch.Tensor,
page_size: int,
kv_len: int,
mask: torch.Tensor | None = None,
is_causal: bool = False,
) -> torch.Tensor:
"""Paged GQA decode attention (q_len == 1, direct page-table access).
Args:
q: [batch, 1, n_heads, head_dim] (blhd, bf16)
page_table: [batch, max_pages] (int64)
k_cache: [n_pages, page_size, n_kv_heads, head_dim] (bf16)
v_cache: same as k_cache
page_size: tokens per page
kv_len: actual sequence length per request
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_paged_decode")
causal_offset = (kv_len - 1) if is_causal else -1
return _modules["attn_paged_decode"].attn_paged_decode(
q,
page_table,
k_cache,
v_cache,
page_size,
kv_len,
mask=mask,
causal_offset=causal_offset,
layout=1,
)
-298
View File
@@ -1,298 +0,0 @@
"""GQA attention wrapper functions — one entry point per compiled kernel.
Each wrapper dispatches to its CUDA kernel (loaded in ``loader.py``) when
available, otherwise falls back to ``torch`` SDPA.
Interface (all functions):
causal_offset: -1 = non-causal; >=0 = absolute position of first Q token
mask: 2D [batch, kv_len] or 3D [batch, q_len, kv_len] (bool)
scale: 0.0 = auto (1/sqrt(head_dim)); >0 = explicit
layout: "bhld" (default) or "blhd"
Add new kernel wrappers here; split into per-variant files only if this file
grows large.
"""
import math
import torch
import torch.nn.functional as F
from astrai.extension.loader import _available, _modules
_LAYOUT_CODES: dict[str, int] = {"bhld": 0, "blhd": 1}
def _parse_layout(layout: str | int) -> int:
if isinstance(layout, int):
return layout
code = _LAYOUT_CODES.get(layout.lower())
if code is None:
raise ValueError(
f"unknown layout '{layout}', expected one of {list(_LAYOUT_CODES)}"
)
return code
def _to_bhld(t: torch.Tensor, layout: int) -> torch.Tensor:
"""Normalize to b h l d view. Zero-copy transpose if layout==1 (b l h d)."""
if layout == 1:
return t.transpose(1, 2)
return t
def _expand_kv_heads(
k: torch.Tensor, v: torch.Tensor, q_head: int
) -> tuple[torch.Tensor, torch.Tensor]:
"""Expand K/V heads to match Q heads for GQA fallback."""
kv_head = k.size(1)
if kv_head == q_head:
return k, v
group = q_head // kv_head
k = k.repeat_interleave(group, dim=1)
v = v.repeat_interleave(group, dim=1)
return k, v
def _build_attn_mask(
q: torch.Tensor,
k: torch.Tensor,
mask: torch.Tensor | None,
causal_offset: int,
scale: float,
) -> tuple[torch.Tensor | None, float]:
"""Build SDPA-compatible attn_mask + resolved scale.
q and k must already be in b h l d layout.
Causal and mask can coexist: causal sets -inf above the diagonal, mask
sets -inf for padded positions. Both are OR'd into a single bool mask.
"""
q_len = q.size(2)
kv_len = k.size(2)
head_dim = q.size(3)
resolved_scale = scale if scale and scale > 0 else 1.0 / math.sqrt(head_dim)
attn_mask = None
if mask is not None:
if mask.dim() == 2:
# [batch, kv_len] → [batch, 1, 1, kv_len]
attn_mask = mask[:, None, None, :]
elif mask.dim() == 3:
# [batch, q_len, kv_len] → [batch, 1, q_len, kv_len]
attn_mask = mask[:, None, :, :]
else:
raise ValueError(f"mask must be 2D or 3D, got {mask.dim()}D")
if causal_offset >= 0:
batch = q.size(0)
# q row i attends to kv cols 0..(causal_offset + i)
q_idx = torch.arange(q_len, device=q.device).unsqueeze(1) # [q_len, 1]
kv_idx = torch.arange(kv_len, device=q.device).unsqueeze(0) # [1, kv_len]
causal_bool = kv_idx > (causal_offset + q_idx) # True = masked out
causal_mask = causal_bool.unsqueeze(0).expand(
batch, -1, -1
) # [batch, q_len, kv_len]
causal_mask = causal_mask[:, None, :, :] # [batch, 1, q_len, kv_len]
if attn_mask is not None:
attn_mask = attn_mask | causal_mask
else:
attn_mask = causal_mask
return attn_mask, resolved_scale
def _torch_fallback(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: torch.Tensor | None,
causal_offset: int,
scale: float,
q_layout: int,
kv_layout: int | None = None,
) -> torch.Tensor:
"""Reference attention via ``scaled_dot_product_attention``.
q_layout / kv_layout: 0 = b h l d, 1 = b l h d.
If kv_layout is None, uses q_layout (Q and K/V share the same layout).
"""
if kv_layout is None:
kv_layout = q_layout
q = _to_bhld(q, q_layout)
k = _to_bhld(k, kv_layout)
v = _to_bhld(v, kv_layout)
k, v = _expand_kv_heads(k, v, q.size(1))
attn_mask, resolved_scale = _build_attn_mask(q, k, mask, causal_offset, scale)
out = F.scaled_dot_product_attention(
q, k, v, attn_mask=attn_mask, is_causal=False, scale=resolved_scale
)
# Restore Q's original layout
if q_layout == 1:
out = out.transpose(1, 2)
return out
def _gather_kv_from_pages(
page_table: torch.Tensor,
k_cache: torch.Tensor,
v_cache: torch.Tensor,
page_size: int,
kv_len: int,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Gather contiguous K/V from paged cache for torch SDPA fallback.
Shapes:
page_table : [batch, max_pages] (int64)
k_cache : [n_pages, page_size, n_kv_heads, head_dim]
v_cache : same as k_cache
Returns:
k, v : [batch, kv_len, n_kv_heads, head_dim] (b l h d)
"""
batch, max_pages = page_table.shape
_, ps, n_kv_heads, head_dim = k_cache.shape
if ps != page_size:
raise ValueError(f"k_cache page_size mismatch: {ps} vs {page_size}")
# Vectorized gather: build physical page + offset indices, then advanced-index
positions = torch.arange(kv_len, device=page_table.device)
logical_pages = positions // page_size # [kv_len]
page_offsets = positions % page_size # [kv_len]
phys_pages = page_table[:, logical_pages] # [batch, kv_len]
# k_cache[phys_pages, page_offsets] → [batch, kv_len, n_kv_heads, head_dim] (b l h d)
k = k_cache[phys_pages, page_offsets]
v = v_cache[phys_pages, page_offsets]
return k, v
def attn_decode(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: torch.Tensor | None = None,
causal_offset: int = -1,
scale: float = 0.0,
layout: str = "bhld",
) -> torch.Tensor:
li = _parse_layout(layout)
if _available["attn_decode"]:
return _modules["attn_decode"].attn_decode(
q,
k,
v,
mask=mask,
causal_offset=causal_offset,
scale=scale,
layout=li,
)
return _torch_fallback(q, k, v, mask, causal_offset, scale, q_layout=li)
def attn_prefill(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: torch.Tensor | None = None,
causal_offset: int = -1,
scale: float = 0.0,
layout: str = "bhld",
) -> torch.Tensor:
li = _parse_layout(layout)
if _available["attn_prefill"]:
return _modules["attn_prefill"].attn_prefill(
q,
k,
v,
mask=mask,
causal_offset=causal_offset,
scale=scale,
layout=li,
)
return _torch_fallback(q, k, v, mask, causal_offset, scale, q_layout=li)
def attn_paged_decode(
q: torch.Tensor,
page_table: torch.Tensor,
k_cache: torch.Tensor,
v_cache: torch.Tensor,
page_size: int,
kv_len: int,
mask: torch.Tensor | None = None,
causal_offset: int = -1,
scale: float = 0.0,
layout: str = "bhld",
) -> torch.Tensor:
li = _parse_layout(layout)
if _available["attn_paged_decode"]:
return _modules["attn_paged_decode"].attn_paged_decode(
q,
page_table,
k_cache,
v_cache,
page_size,
kv_len,
mask=mask,
causal_offset=causal_offset,
scale=scale,
layout=li,
)
# Gathered K/V are always b l h d
k, v = _gather_kv_from_pages(page_table, k_cache, v_cache, page_size, kv_len)
return _torch_fallback(
q, k, v, mask, causal_offset, scale, q_layout=li, kv_layout=1
)
def attention(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
mask: torch.Tensor | None = None,
causal_offset: int = -1,
scale: float = 0.0,
layout: str = "bhld",
) -> torch.Tensor:
"""Dispatch to decode or prefill attention based on the query length.
A query length of one is the decode case; longer queries use prefill.
The paged-cache decode path cannot be selected here because its page-table
arguments are not part of this interface.
"""
li = _parse_layout(layout)
if q.ndim not in (2, 3, 4) or k.ndim != q.ndim or v.ndim != q.ndim:
raise ValueError(
"q, k, and v must all have the same rank in {2, 3, 4}, "
f"got {q.ndim}D, {k.ndim}D, {v.ndim}D"
)
if k.shape != v.shape:
raise ValueError(
f"k and v must have the same shape, got {k.shape} and {v.shape}"
)
original_ndim = q.ndim
if original_ndim == 2:
# [L, D] -> [1, 1, L, D] or [1, L, 1, D]
q = q.unsqueeze(0).unsqueeze(1 if li == 0 else 2)
k = k.unsqueeze(0).unsqueeze(1 if li == 0 else 2)
v = v.unsqueeze(0).unsqueeze(1 if li == 0 else 2)
elif original_ndim == 3:
# [B, L, D] -> single-head 4D input.
q = q.unsqueeze(1 if li == 0 else 2)
k = k.unsqueeze(1 if li == 0 else 2)
v = v.unsqueeze(1 if li == 0 else 2)
q_len = q.size(2 if li == 0 else 1)
if q_len == 1:
out = attn_decode(q, k, v, mask, causal_offset, scale, layout)
else:
out = attn_prefill(q, k, v, mask, causal_offset, scale, layout)
if original_ndim == 2:
return out.squeeze(0).squeeze(0 if li == 0 else 1)
if original_ndim == 3:
return out.squeeze(1 if li == 0 else 2)
return out