diff --git a/astrai/extension/__init__.py b/astrai/extension/__init__.py index 7fba00d..99bb03b 100644 --- a/astrai/extension/__init__.py +++ b/astrai/extension/__init__.py @@ -19,6 +19,7 @@ from astrai.extension.attention_backend import ( ATTN_BACKEND, AttentionBackend, CudaBackend, + FlashAttnBackend, TorchNativeBackend, attention, attn_backend, @@ -38,6 +39,7 @@ __all__ = [ "AttentionBackend", "CudaBackend", "TorchNativeBackend", + "FlashAttnBackend", "TensorLayout", "attention", "attn_backend", diff --git a/astrai/extension/attention_backend.py b/astrai/extension/attention_backend.py index d4221c2..eec63cc 100644 --- a/astrai/extension/attention_backend.py +++ b/astrai/extension/attention_backend.py @@ -30,6 +30,8 @@ Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]`` import contextvars import enum +import importlib +import threading from abc import ABC, abstractmethod from contextlib import contextmanager from typing import Optional, Union @@ -48,12 +50,87 @@ _current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.Conte "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): """Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``.""" TORCH_NATIVE = "torch_native" CUDA = "cuda" + FLASH = "flash" def get_backend() -> "AttentionBackend": @@ -294,11 +371,13 @@ class TorchNativeBackend(AttentionBackend): k = repeat_kv(k, n_rep) v = repeat_kv(v, n_rep) - q = q.permute(0, 2, 1, 3) - k = k.permute(0, 2, 1, 3) - v = v.permute(0, 2, 1, 3) - - out = F.scaled_dot_product_attention(q, k, v, attn_mask, is_causal=is_causal) + out = F.scaled_dot_product_attention( + q.permute(0, 2, 1, 3), + k.permute(0, 2, 1, 3), + v.permute(0, 2, 1, 3), + attn_mask, + is_causal=is_causal, + ) out = out.permute(0, 2, 1, 3).contiguous().flatten(2) return out @@ -395,7 +474,93 @@ class CudaBackend(AttentionBackend): 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]] = { ATTN_BACKEND.TORCH_NATIVE: TorchNativeBackend, ATTN_BACKEND.CUDA: CudaBackend, + ATTN_BACKEND.FLASH: FlashAttnBackend, } diff --git a/pyproject.toml b/pyproject.toml index 4ef22dc..856a77e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -32,6 +32,7 @@ urls = { Homepage = "https://github.com/ViperEkura/AstrAI" } [project.optional-dependencies] dev = ["pytest==9.0.2", "ruff", "httpx2"] +flash = ["flash-attn>=2.6"] [tool.setuptools.packages.find] where = ["."]