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:
@@ -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",
|
||||||
|
|||||||
@@ -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,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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 = ["."]
|
||||||
|
|||||||
Reference in New Issue
Block a user