feat: add optional FlashAttention (FA2/FA3) backend

- add FlashAttnBackend (ATTN_BACKEND.FLASH) using flash_attn_func with KV-cache gather + GQA, mirroring TorchNativeBackend
- add flash_attn_available() probe gated on compute capability plus a real-kernel smoke test, cached at first use
- lazy-import flash-attn via importlib so it stays an optional dependency, raising clear errors when unusable
- add 'flash' optional extra (flash-attn>=2.6) and export the new backend
This commit is contained in:
2026-08-05 15:27:26 +08:00
parent 2667b8116d
commit 8c052c99ee
3 changed files with 173 additions and 5 deletions
+2
View File
@@ -19,6 +19,7 @@ from astrai.extension.attention_backend import (
ATTN_BACKEND, ATTN_BACKEND,
AttentionBackend, AttentionBackend,
CudaBackend, CudaBackend,
FlashAttnBackend,
TorchNativeBackend, TorchNativeBackend,
attention, attention,
attn_backend, attn_backend,
@@ -38,6 +39,7 @@ __all__ = [
"AttentionBackend", "AttentionBackend",
"CudaBackend", "CudaBackend",
"TorchNativeBackend", "TorchNativeBackend",
"FlashAttnBackend",
"TensorLayout", "TensorLayout",
"attention", "attention",
"attn_backend", "attn_backend",
+170 -5
View File
@@ -30,6 +30,8 @@ Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
import contextvars import contextvars
import enum import enum
import importlib
import threading
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
@@ -48,12 +50,87 @@ _current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.Conte
"attn_backend" "attn_backend"
) )
_lock = threading.Lock()
_flash_available: Optional[bool] = None
def flash_attn_available() -> bool:
"""Return ``True`` if the optional ``flash-attn`` package is usable.
``flash-attn`` is not a hard dependency (declared only as an optional
extra and imported lazily), so this is checked at first use and cached.
The check is stronger than "import works": it also gates on the GPU
compute capability for the installed major version and smoke-tests a
real tiny kernel call, because wheels that import fine can still fail
at the first actual invocation (wrong arch build, torch mismatch, or a
missing ``flash_attn_func`` entry point). It never raises.
"""
global _flash_available
if _flash_available is None:
with _lock:
if _flash_available is None:
_flash_available = _flash_attn_check()
return _flash_available
_flash_attn_module = None
_flash_attn_import_tried = False
def _get_flash_attn():
"""Lazily import and cache the optional ``flash_attn`` module.
Uses ``importlib.import_module`` so no static import binds the name when
the package is absent. Returns the module object, or ``None`` if the
package is not installed or cannot be imported. Never raises.
"""
global _flash_attn_module, _flash_attn_import_tried
if not _flash_attn_import_tried:
_flash_attn_import_tried = True
try:
_flash_attn_module = importlib.import_module("flash_attn")
except Exception:
_flash_attn_module = None
return _flash_attn_module
def _flash_attn_check() -> bool:
if not torch.cuda.is_available():
return False
fa = _get_flash_attn()
if fa is None:
return False
# version + compute-capability gate:
# FlashAttention-2 kernels need sm_70+; FlashAttention-3 (tcgen05,
# sm_90/sm_100) needs sm_90+.
try:
major = int(fa.__version__.split(".")[0])
cc = torch.cuda.get_device_capability()
cc_num = cc[0] * 10 + cc[1]
except Exception:
major, cc_num = 0, 0
if (major >= 3 and cc_num < 90) or (major < 3 and 0 < cc_num < 70):
return False
# smoke-test the real kernel: a wheel that imports but was built for a
# different arch/torch fails here instead of at the first real forward.
try:
if not hasattr(fa, "flash_attn_func"):
return False
x = torch.zeros(1, 1, 1, 64, device="cuda", dtype=torch.bfloat16)
out = fa.flash_attn_func(x, x, x, causal=True)
return bool(torch.isfinite(out).all().item())
except Exception:
return False
class ATTN_BACKEND(enum.Enum): 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" CUDA = "cuda"
FLASH = "flash"
def get_backend() -> "AttentionBackend": def get_backend() -> "AttentionBackend":
@@ -294,11 +371,13 @@ class TorchNativeBackend(AttentionBackend):
k = repeat_kv(k, n_rep) k = repeat_kv(k, n_rep)
v = repeat_kv(v, n_rep) v = repeat_kv(v, n_rep)
q = q.permute(0, 2, 1, 3) out = F.scaled_dot_product_attention(
k = k.permute(0, 2, 1, 3) q.permute(0, 2, 1, 3),
v = v.permute(0, 2, 1, 3) k.permute(0, 2, 1, 3),
v.permute(0, 2, 1, 3),
out = F.scaled_dot_product_attention(q, k, v, attn_mask, is_causal=is_causal) attn_mask,
is_causal=is_causal,
)
out = out.permute(0, 2, 1, 3).contiguous().flatten(2) out = out.permute(0, 2, 1, 3).contiguous().flatten(2)
return out return out
@@ -395,7 +474,93 @@ class CudaBackend(AttentionBackend):
return out.reshape(b, q_len, q.size(2), q.size(3)).flatten(2) return out.reshape(b, q_len, q.size(2), q.size(3)).flatten(2)
class FlashAttnBackend(AttentionBackend):
"""FlashAttention (FA2/FA3) backend via the optional ``flash-attn`` package.
Uses the general ``flash_attn_func`` entry point for both prefill and
single-token decode, mirroring ``TorchNativeBackend``'s KV-cache gather.
This backend only does flash attention — inputs ``flash-attn`` cannot
express (missing package, custom attention mask, fp32, unsupported
head_dim) raise a clear error instead of silently falling back to torch.
For a torch fallback, select ``TorchNativeBackend`` instead.
"""
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:
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
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:
return self._forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
def _forward(
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 not None:
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.max_len
indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
if q.size(1) == 1 and attn_mask is not None and attn_mask.dim() == 4:
pos_mask = attn_mask[:, 0, 0]
else:
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 = kv_cache.k_buffer[layer_id, indices]
v = kv_cache.v_buffer[layer_id, indices]
n_rep = q.size(2) // k.size(2)
if n_rep > 1:
k = repeat_kv(k, n_rep)
v = repeat_kv(v, n_rep)
if attn_mask is not None and not is_causal:
raise ValueError(
"FlashAttnBackend does not support a custom attention mask; "
"use a causal mask or select TorchNativeBackend."
)
fa = _get_flash_attn()
if fa is None:
raise RuntimeError(
"FlashAttnBackend requires the optional 'flash-attn' package. "
"Install with `pip install flash-attn`."
)
out = fa.flash_attn_func(
q.contiguous(), k.contiguous(), v.contiguous(), causal=is_causal
)
return out.contiguous().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, ATTN_BACKEND.CUDA: CudaBackend,
ATTN_BACKEND.FLASH: FlashAttnBackend,
} }
+1
View File
@@ -32,6 +32,7 @@ urls = { Homepage = "https://github.com/ViperEkura/AstrAI" }
[project.optional-dependencies] [project.optional-dependencies]
dev = ["pytest==9.0.2", "ruff", "httpx2"] dev = ["pytest==9.0.2", "ruff", "httpx2"]
flash = ["flash-attn>=2.6"]
[tool.setuptools.packages.find] [tool.setuptools.packages.find]
where = ["."] where = ["."]