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:
@@ -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",
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -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(
|
|
||||||
q, k, v, kv_cache, layer_id, n_rep, attn_mask, is_causal, scale
|
|
||||||
)
|
)
|
||||||
|
return self.forward_extend(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
|
||||||
|
|
||||||
@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
|
||||||
|
|
||||||
|
|||||||
@@ -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,
|
||||||
|
)
|
||||||
@@ -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
|
|
||||||
Reference in New Issue
Block a user