feat: add CudaBackend and rename to fwd_decode/fwd_prefill
- CudaBackend: paged decode via attn_paged_decode, prefill via attn_prefill - Decode uses req_to_token as page_table with page_size=1 - Falls back to TorchNativeBackend when kernel unavailable - Rename forward_decode/forward_extend to fwd_decode/fwd_prefill - Register ATTN_BACKEND.CUDA in _BACKEND_REGISTRY
This commit is contained in:
@@ -6,17 +6,19 @@ Public API:
|
|||||||
- ``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
|
- ``AttentionBackend`` — ABC for attention computation strategies
|
||||||
- ``TorchNativeBackend`` — default SDPA backend with KV cache I/O
|
- ``TorchNativeBackend`` — default SDPA backend with KV cache I/O
|
||||||
|
- ``CudaBackend`` — CUDA kernel backend with paged decode + prefill
|
||||||
|
|
||||||
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
||||||
(blhd). Scale is always ``1/sqrt(head_dim)``.
|
(blhd). Scale is always ``1/sqrt(head_dim)``.
|
||||||
|
|
||||||
Each wrapper dispatches to its compiled CUDA kernel (``astrai.extension.attn_*``)
|
Each wrapper calls its compiled CUDA kernel directly. Fallback to torch
|
||||||
when available, otherwise falls back to ``torch.nn.functional.scaled_dot_product_attention``.
|
SDPA is handled by the attention backend, not the wrapper functions.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from astrai.extension.attention_backend import (
|
from astrai.extension.attention_backend import (
|
||||||
ATTN_BACKEND,
|
ATTN_BACKEND,
|
||||||
AttentionBackend,
|
AttentionBackend,
|
||||||
|
CudaBackend,
|
||||||
TorchNativeBackend,
|
TorchNativeBackend,
|
||||||
attn_backend,
|
attn_backend,
|
||||||
get_backend,
|
get_backend,
|
||||||
@@ -31,6 +33,7 @@ from astrai.extension.loader import KERNEL_NAMES, is_available
|
|||||||
__all__ = [
|
__all__ = [
|
||||||
"ATTN_BACKEND",
|
"ATTN_BACKEND",
|
||||||
"AttentionBackend",
|
"AttentionBackend",
|
||||||
|
"CudaBackend",
|
||||||
"TorchNativeBackend",
|
"TorchNativeBackend",
|
||||||
"attn_backend",
|
"attn_backend",
|
||||||
"get_backend",
|
"get_backend",
|
||||||
|
|||||||
@@ -39,6 +39,8 @@ import torch
|
|||||||
import torch.nn.functional as F
|
import torch.nn.functional as F
|
||||||
from torch import Tensor
|
from torch import Tensor
|
||||||
|
|
||||||
|
from astrai.extension.attention_ops import attn_paged_decode, attn_prefill
|
||||||
|
from astrai.extension.loader import is_available
|
||||||
from astrai.inference.core.cache import KVCache
|
from astrai.inference.core.cache import KVCache
|
||||||
|
|
||||||
_current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.ContextVar(
|
_current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.ContextVar(
|
||||||
@@ -50,6 +52,7 @@ class ATTN_BACKEND(enum.Enum):
|
|||||||
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
|
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
|
||||||
|
|
||||||
TORCH_NATIVE = "torch_native"
|
TORCH_NATIVE = "torch_native"
|
||||||
|
CUDA = "cuda"
|
||||||
|
|
||||||
|
|
||||||
def get_backend() -> "AttentionBackend":
|
def get_backend() -> "AttentionBackend":
|
||||||
@@ -113,8 +116,8 @@ def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
|
|||||||
class AttentionBackend(ABC):
|
class AttentionBackend(ABC):
|
||||||
"""Abstract base for attention computation strategies.
|
"""Abstract base for attention computation strategies.
|
||||||
|
|
||||||
Subclasses implement ``forward_decode`` (q_len == 1, with cache) and
|
Subclasses implement ``fwd_decode`` (q_len == 1, with cache) and
|
||||||
``forward_extend`` (q_len > 1, with or without cache). The public
|
``fwd_prefill`` (q_len > 1, with or without cache). The public
|
||||||
``forward`` method dispatches based on q_len.
|
``forward`` method dispatches based on q_len.
|
||||||
|
|
||||||
Three equivalent ways to activate a backend::
|
Three equivalent ways to activate a backend::
|
||||||
@@ -159,13 +162,11 @@ class AttentionBackend(ABC):
|
|||||||
[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.fwd_decode(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||||
q, k, v, kv_cache, layer_id, attn_mask, is_causal
|
return self.fwd_prefill(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||||
)
|
|
||||||
return self.forward_extend(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def forward_decode(
|
def fwd_decode(
|
||||||
self,
|
self,
|
||||||
q: Tensor,
|
q: Tensor,
|
||||||
k: Tensor,
|
k: Tensor,
|
||||||
@@ -178,7 +179,7 @@ class AttentionBackend(ABC):
|
|||||||
"""Single-token decode with KV cache."""
|
"""Single-token decode with KV cache."""
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def forward_extend(
|
def fwd_prefill(
|
||||||
self,
|
self,
|
||||||
q: Tensor,
|
q: Tensor,
|
||||||
k: Tensor,
|
k: Tensor,
|
||||||
@@ -202,7 +203,7 @@ class TorchNativeBackend(AttentionBackend):
|
|||||||
runs SDPA directly on the projected q/k/v.
|
runs SDPA directly on the projected q/k/v.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def forward_decode(
|
def fwd_decode(
|
||||||
self,
|
self,
|
||||||
q: Tensor,
|
q: Tensor,
|
||||||
k: Tensor,
|
k: Tensor,
|
||||||
@@ -214,7 +215,7 @@ class TorchNativeBackend(AttentionBackend):
|
|||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||||
|
|
||||||
def forward_extend(
|
def fwd_prefill(
|
||||||
self,
|
self,
|
||||||
q: Tensor,
|
q: Tensor,
|
||||||
k: Tensor,
|
k: Tensor,
|
||||||
@@ -266,6 +267,107 @@ class TorchNativeBackend(AttentionBackend):
|
|||||||
|
|
||||||
_default_backend = TorchNativeBackend()
|
_default_backend = TorchNativeBackend()
|
||||||
|
|
||||||
|
|
||||||
|
class CudaBackend(AttentionBackend):
|
||||||
|
"""CUDA kernel backend with direct KV cache access.
|
||||||
|
|
||||||
|
Decode path: writes K/V to cache, then calls ``attn_paged_decode``
|
||||||
|
with ``page_size=1`` (each token slot is a single-token "page").
|
||||||
|
The ``req_to_token`` table serves directly as the page table.
|
||||||
|
|
||||||
|
Prefill path: writes K/V to cache, gathers full-sequence K/V via
|
||||||
|
indirect indexing (same as TorchNativeBackend), then calls
|
||||||
|
``attn_prefill``.
|
||||||
|
|
||||||
|
Training path (``kv_cache is None``): calls ``attn_prefill`` directly
|
||||||
|
on the projected q/k/v.
|
||||||
|
|
||||||
|
Falls back to ``TorchNativeBackend`` for any path where the
|
||||||
|
corresponding CUDA kernel is not available.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self._fallback = TorchNativeBackend()
|
||||||
|
|
||||||
|
def fwd_decode(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
if kv_cache is None or not is_available("attn_paged_decode"):
|
||||||
|
return self._fallback.fwd_decode(
|
||||||
|
q, k, v, kv_cache, layer_id, attn_mask, is_causal
|
||||||
|
)
|
||||||
|
|
||||||
|
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||||
|
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||||
|
|
||||||
|
max_len = kv_cache.seq_lens.max().item()
|
||||||
|
|
||||||
|
page_table = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
||||||
|
|
||||||
|
k_cache = kv_cache.k_buffer[layer_id].unsqueeze(1)
|
||||||
|
v_cache = kv_cache.v_buffer[layer_id].unsqueeze(1)
|
||||||
|
|
||||||
|
out = attn_paged_decode(
|
||||||
|
q,
|
||||||
|
page_table,
|
||||||
|
k_cache,
|
||||||
|
v_cache,
|
||||||
|
page_size=1,
|
||||||
|
kv_len=max_len,
|
||||||
|
mask=None,
|
||||||
|
is_causal=is_causal,
|
||||||
|
)
|
||||||
|
|
||||||
|
out = out.flatten(2)
|
||||||
|
return out
|
||||||
|
|
||||||
|
def fwd_prefill(
|
||||||
|
self,
|
||||||
|
q: Tensor,
|
||||||
|
k: Tensor,
|
||||||
|
v: Tensor,
|
||||||
|
kv_cache: Optional[KVCache],
|
||||||
|
layer_id: int,
|
||||||
|
attn_mask: Optional[Tensor] = None,
|
||||||
|
is_causal: bool = False,
|
||||||
|
) -> Tensor:
|
||||||
|
if kv_cache is None:
|
||||||
|
if is_available("attn_prefill"):
|
||||||
|
out = attn_prefill(q, k, v, mask=attn_mask, is_causal=is_causal)
|
||||||
|
return out.flatten(2)
|
||||||
|
return self._fallback.fwd_prefill(
|
||||||
|
q, k, v, kv_cache, layer_id, attn_mask, is_causal
|
||||||
|
)
|
||||||
|
|
||||||
|
if not is_available("attn_prefill"):
|
||||||
|
return self._fallback.fwd_prefill(
|
||||||
|
q, k, v, kv_cache, layer_id, attn_mask, is_causal
|
||||||
|
)
|
||||||
|
|
||||||
|
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
||||||
|
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
||||||
|
|
||||||
|
max_len = kv_cache.seq_lens.max()
|
||||||
|
indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
||||||
|
pos_mask = (
|
||||||
|
torch.arange(max_len, device=q.device)[None, :] < kv_cache.seq_lens[:, None]
|
||||||
|
)
|
||||||
|
indices = torch.where(pos_mask, indices, torch.zeros_like(indices))
|
||||||
|
k_full = kv_cache.k_buffer[layer_id, indices]
|
||||||
|
v_full = kv_cache.v_buffer[layer_id, indices]
|
||||||
|
|
||||||
|
out = attn_prefill(q, k_full, v_full, mask=attn_mask, is_causal=is_causal)
|
||||||
|
return out.flatten(2)
|
||||||
|
|
||||||
|
|
||||||
_BACKEND_REGISTRY: dict[ATTN_BACKEND, type[AttentionBackend]] = {
|
_BACKEND_REGISTRY: dict[ATTN_BACKEND, type[AttentionBackend]] = {
|
||||||
ATTN_BACKEND.TORCH_NATIVE: TorchNativeBackend,
|
ATTN_BACKEND.TORCH_NATIVE: TorchNativeBackend,
|
||||||
|
ATTN_BACKEND.CUDA: CudaBackend,
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user