diff --git a/astrai/extension/__init__.py b/astrai/extension/__init__.py index 068e0c6..d41f44f 100644 --- a/astrai/extension/__init__.py +++ b/astrai/extension/__init__.py @@ -15,25 +15,25 @@ Each wrapper calls its compiled CUDA kernel directly. Fallback to torch SDPA is handled by the attention backend, not the wrapper functions. """ -from astrai.extension.attention_backend import ( +from astrai.extension.backend import ( ATTN_BACKEND, AttentionBackend, AttentionBackendFactory, CudaBackend, FlashAttnBackend, TorchNativeBackend, + apply_rotary_emb, attention, attn_backend, get_backend, ) -from astrai.extension.attention_ops import ( +from astrai.extension.loader import KERNEL_NAMES, is_available +from astrai.extension.ops import ( TensorLayout, attn_decode, attn_paged_decode, attn_prefill, ) -from astrai.extension.loader import KERNEL_NAMES, is_available -from astrai.extension.rotary_backend import apply_rotary_emb __all__ = [ "ATTN_BACKEND", diff --git a/astrai/extension/backend/__init__.py b/astrai/extension/backend/__init__.py new file mode 100644 index 0000000..189f8e8 --- /dev/null +++ b/astrai/extension/backend/__init__.py @@ -0,0 +1,27 @@ +"""Backend selection, fallbacks, and execution policies.""" + +from astrai.extension.backend.attention import ( + ATTN_BACKEND, + AttentionBackend, + AttentionBackendFactory, + CudaBackend, + FlashAttnBackend, + TorchNativeBackend, + attention, + attn_backend, + get_backend, +) +from astrai.extension.backend.rotary import apply_rotary_emb + +__all__ = [ + "ATTN_BACKEND", + "AttentionBackend", + "AttentionBackendFactory", + "CudaBackend", + "FlashAttnBackend", + "TorchNativeBackend", + "apply_rotary_emb", + "attention", + "attn_backend", + "get_backend", +] diff --git a/astrai/extension/attention_backend.py b/astrai/extension/backend/attention.py similarity index 98% rename from astrai/extension/attention_backend.py rename to astrai/extension/backend/attention.py index fa5c0c0..7921bd0 100644 --- a/astrai/extension/attention_backend.py +++ b/astrai/extension/backend/attention.py @@ -32,7 +32,6 @@ Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]`` import contextvars import enum import functools -import importlib import os import threading from abc import ABC, abstractmethod @@ -43,13 +42,18 @@ import torch import torch.nn.functional as F from torch import Tensor -from astrai.extension.attention_ops import ( +from astrai.extension.loader import is_available +from astrai.extension.ops.attention import ( attn_paged_decode, attn_paged_prefill, ) -from astrai.extension.loader import is_available from astrai.factory import BaseFactory +try: + import flash_attn as _flash_attn +except Exception: + _flash_attn = None + if TYPE_CHECKING: from astrai.inference.cache import KVCache @@ -67,7 +71,7 @@ _current_backend: contextvars.ContextVar[Optional["AttentionBackend"]] = ( def flash_attn_available() -> bool: if not torch.cuda.is_available(): return False - fa = _get_flash_attn() + fa = _flash_attn if fa is None: return False @@ -90,14 +94,6 @@ def flash_attn_available() -> bool: return False -@functools.lru_cache(maxsize=1) -def _get_flash_attn(): - try: - return importlib.import_module("flash_attn") - except Exception: - return None - - class ATTN_BACKEND(enum.Enum): """Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``.""" @@ -145,7 +141,7 @@ def _backend_supports( if q.dtype not in (torch.float16, torch.bfloat16): return False if fwd is not None: - return q.ndim == 3 and hasattr(_get_flash_attn(), "flash_attn_varlen_func") + return q.ndim == 3 and hasattr(_flash_attn, "flash_attn_varlen_func") if attn_mask is None or is_causal: return True return attn_mask.dim() == 4 @@ -680,7 +676,7 @@ class FlashAttnBackend(AttentionBackend): "FlashAttnBackend does not support a custom attention mask; " "use a causal mask or select TorchNativeBackend." ) - fa = _get_flash_attn() + fa = _flash_attn if fa is None: raise RuntimeError( "FlashAttnBackend requires the optional 'flash-attn' package. " @@ -702,7 +698,7 @@ class FlashAttnBackend(AttentionBackend): kv_cache: "KVCache", layer_id: int, ) -> Tensor: - fa = _get_flash_attn() + fa = _flash_attn if fa is None or not hasattr(fa, "flash_attn_varlen_func"): raise RuntimeError("packed inference requires flash_attn_varlen_func") kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k diff --git a/astrai/extension/rotary_backend.py b/astrai/extension/backend/rotary.py similarity index 95% rename from astrai/extension/rotary_backend.py rename to astrai/extension/backend/rotary.py index 91bd2ef..f934f0f 100644 --- a/astrai/extension/rotary_backend.py +++ b/astrai/extension/backend/rotary.py @@ -11,6 +11,7 @@ import torch from torch import Tensor from astrai.extension.loader import is_available +from astrai.extension.ops.rotary import rotary_emb as _cuda_rotary _cache = {"available": None} @@ -48,7 +49,5 @@ def apply_rotary_emb(x: Tensor, freqs_cis: Tensor) -> Tensor: and x.is_cuda and x.dtype == torch.bfloat16 ): - from astrai.extension.rotary_ops import rotary_emb as _cuda_rotary - return _cuda_rotary(x, freqs_cis) return _torch_apply(x, freqs_cis) diff --git a/astrai/extension/fp8.py b/astrai/extension/fp8.py index 74d4ec8..db1606d 100644 --- a/astrai/extension/fp8.py +++ b/astrai/extension/fp8.py @@ -1,8 +1,8 @@ """FP8 training: scaling state and aten::linear dispatch. -Layered (see also ``fp8_ops.py`` for the CUDA interface adapter): +Layered (see also ``ops/fp8.py`` for the CUDA interface adapter): -1. Kernel interface: "fp8_ops" — the only module touching the pybind. +1. Kernel interface: ``ops.fp8`` - the only module touching the pybind. 2. Training state (this module): per-tensor scales, amax history, delayed scaling, and the ``fp8_autocast`` context (TE-style, like ``torch.autocast``). @@ -25,7 +25,7 @@ from contextlib import contextmanager import torch from torch.library import Library -from astrai.extension.fp8_ops import ( +from astrai.extension.ops.fp8 import ( linear_backward_scaled, linear_forward_scaled, ) diff --git a/astrai/extension/ops/__init__.py b/astrai/extension/ops/__init__.py new file mode 100644 index 0000000..ac0009a --- /dev/null +++ b/astrai/extension/ops/__init__.py @@ -0,0 +1,19 @@ +"""Stateless wrappers around compiled extension kernels.""" + +from astrai.extension.ops.attention import ( + TensorLayout, + attn_decode, + attn_paged_decode, + attn_paged_prefill, + attn_prefill, +) +from astrai.extension.ops.rotary import rotary_emb + +__all__ = [ + "TensorLayout", + "attn_decode", + "attn_paged_decode", + "attn_paged_prefill", + "attn_prefill", + "rotary_emb", +] diff --git a/astrai/extension/attention_ops.py b/astrai/extension/ops/attention.py similarity index 98% rename from astrai/extension/attention_ops.py rename to astrai/extension/ops/attention.py index 05cc8d8..47a8d55 100644 --- a/astrai/extension/attention_ops.py +++ b/astrai/extension/ops/attention.py @@ -1,4 +1,4 @@ -"""Attention kernel wrapper functions — one entry point per compiled kernel. +"""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 diff --git a/astrai/extension/fp8_ops.py b/astrai/extension/ops/fp8.py similarity index 100% rename from astrai/extension/fp8_ops.py rename to astrai/extension/ops/fp8.py diff --git a/astrai/extension/rotary_ops.py b/astrai/extension/ops/rotary.py similarity index 94% rename from astrai/extension/rotary_ops.py rename to astrai/extension/ops/rotary.py index 37d2a90..da0f276 100644 --- a/astrai/extension/rotary_ops.py +++ b/astrai/extension/ops/rotary.py @@ -2,7 +2,7 @@ Calls the compiled CUDA kernel directly. If the kernel is not available, raises ``RuntimeError``. Fallback to torch complex multiply is the -responsibility of ``astrai.extension.rotary_backend.apply_rotary_emb``. +responsibility of ``astrai.extension.backend.rotary.apply_rotary_emb``. Layout: x is packed [tokens, n_heads, head_dim] or dense [batch, seq_len, n_heads, head_dim]. ``freqs_cis`` has matching token axes. diff --git a/astrai/inference/runtime/executor.py b/astrai/inference/runtime/executor.py index 0b55cf9..e8c0584 100644 --- a/astrai/inference/runtime/executor.py +++ b/astrai/inference/runtime/executor.py @@ -7,7 +7,7 @@ from typing import List, Optional import torch from torch import Tensor -from astrai.extension.attention_backend import ( +from astrai.extension.backend.attention import ( CudaBackend, get_backend, ) diff --git a/astrai/model/components/__init__.py b/astrai/model/components/__init__.py index 6205674..9f998be 100644 --- a/astrai/model/components/__init__.py +++ b/astrai/model/components/__init__.py @@ -1,4 +1,4 @@ -from astrai.extension.rotary_backend import apply_rotary_emb +from astrai.extension.backend.rotary import apply_rotary_emb from astrai.model.components.attention import GQA, MLA from astrai.model.components.decoder_block import DecoderBlock from astrai.model.components.embedding import Embedding diff --git a/astrai/model/components/attention.py b/astrai/model/components/attention.py index 831323f..f3d300f 100644 --- a/astrai/model/components/attention.py +++ b/astrai/model/components/attention.py @@ -5,8 +5,7 @@ import torch.nn as nn import torch.nn.functional as F from torch import Tensor -from astrai.extension import attention -from astrai.extension.rotary_backend import apply_rotary_emb +from astrai.extension.backend import apply_rotary_emb, attention from astrai.factory import BaseFactory from astrai.inference.cache import KVCache from astrai.model.components.linear import Linear diff --git a/docs/developer/architecture.md b/docs/developer/architecture.md index 446e819..18f23ce 100644 --- a/docs/developer/architecture.md +++ b/docs/developer/architecture.md @@ -1456,7 +1456,7 @@ classDiagram | **Context** | `TrainContext` | Unified training state bag | | **Object Pool** | `Allocator`, `PagePool` | Page-based KV cache with LRU eviction | | **Strategy (Attention)** | `AttentionBackend`, `CudaBackend`, `FlashAttnBackend`, `TorchNativeBackend` | Attention computation backend switching via context manager | -| **Auto-dispatch (Rotary)** | `apply_rotary_emb`, `rotary_backend.py`, `rotary_ops.py` | Rotary embedding CUDA kernel auto-dispatch with torch fallback | +| **Auto-dispatch (Rotary)** | `apply_rotary_emb`, `backend/rotary.py`, `ops/rotary.py` | Rotary embedding CUDA kernel auto-dispatch with torch fallback | | **Executor** | `BaseExecutor`, `NoneExecutor`, `DDPExecutor`, `FSDPExecutor` | Gradient accumulation & model distribution | | **Storage** | `Store`, `MmapStore`, `JsonlStore` | Format-agnostic data access with multi-segment support | | **Producer-Consumer** | `InferenceScheduler`, `Task`, queues | Continuous batching | diff --git a/docs/developer/cuda_kernels.md b/docs/developer/cuda_kernels.md index 88a6f6b..93846c2 100644 --- a/docs/developer/cuda_kernels.md +++ b/docs/developer/cuda_kernels.md @@ -30,7 +30,7 @@ The `rotary_emb` kernel (`csrc/kernels/rotary_emb.cu`) fuses cos/sin lookup and - One thread per (head, dim-pair), vectorized `__nv_bfloat162` load/store - f32 cos/sin input, bf16 compute and output - 256-thread blocks, grid-stride loop -- Auto-dispatched via `apply_rotary_emb` in `astrai/extension/rotary_backend.py` (CUDA when available + inference mode, else torch complex-multiply fallback) +- Auto-dispatched via `apply_rotary_emb` in `astrai/extension/backend/rotary.py` (CUDA when available + inference mode, else torch complex-multiply fallback) - No context-manager backend needed — rotary is backend-agnostic, both attention backends benefit Standalone benchmark vs torch complex-multiply (48 calls = 24 layers × q+k): 6-9x faster, max diff 0 (decode) to 3e-2 (large prefill, bf16). @@ -83,7 +83,7 @@ Each kernel in `astrai/extension/lib` is compiled as an independent pybind11 mod ## Attention Backend -`astrai/extension/attention_backend.py` provides the backend abstraction: +`astrai/extension/backend/attention.py` provides the backend abstraction: - **`AttentionBackend`** (ABC): `fwd_decode` / `fwd_prefill` abstract methods, `forward` dispatches by q_len - **`CudaBackend`**: CUDA kernel dispatch — decode via `attn_paged_decode` (page_size=1), prefill via `attn_paged_prefill` (ragged batch, `qo_indptr` + `kv_indptr`). Default on GPU. @@ -106,7 +106,7 @@ with attn_backend(ATTN_BACKEND.CUDA): ### Rotary Backend -`astrai/extension/rotary_backend.py` provides `apply_rotary_emb(x, (cos, sin))` with auto-dispatch: +`astrai/extension/backend/rotary.py` provides `apply_rotary_emb(x, (cos, sin))` with auto-dispatch: - **CUDA path**: calls `rotary_emb` kernel directly when available, input is bf16 on CUDA, and `torch.is_grad_enabled()` is `False` (inference) - **Torch fallback**: complex multiply (`torch.view_as_complex` → `torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd) or when kernel unavailable @@ -115,9 +115,9 @@ No context-manager switching needed — the dispatch is automatic per call. ## Python Wrappers -`astrai/extension/attention_ops.py` provides Python wrappers for each compiled attention kernel. Each wrapper calls its CUDA kernel directly and raises `RuntimeError` if the `.so` is not available. Fallback to torch SDPA is handled by the attention backend, not the wrapper functions. +`astrai/extension/ops/attention.py` provides Python wrappers for each compiled attention kernel. Each wrapper calls its CUDA kernel directly and raises `RuntimeError` if the `.so` is not available. Fallback to torch SDPA is handled by the attention backend, not the wrapper functions. -`astrai/extension/rotary_ops.py` provides the wrapper for the rotary embedding kernel. Fallback to torch complex multiply is handled by `rotary_backend.py`. +`astrai/extension/ops/rotary.py` provides the wrapper for the rotary embedding kernel. Fallback to torch complex multiply is handled by `backend/rotary.py`. Interface (all functions): ``` diff --git a/docs/developer/internals.md b/docs/developer/internals.md index 329a774..bb28f6c 100644 --- a/docs/developer/internals.md +++ b/docs/developer/internals.md @@ -176,14 +176,14 @@ Three-layer separation (SGLang-inspired): ### Attention Backend -Attention computation is decoupled from the model via `AttentionBackend` ABC (`astrai/extension/attention_backend.py`): +Attention computation is decoupled from the model via `AttentionBackend` ABC (`astrai/extension/backend/attention.py`): - **`CudaBackend`** (default): decode path uses `attn_paged_decode` with `page_size=1` (the `req_to_token` table serves as the page table, each token slot is a single-token "page"); prefill path uses the ragged-batch `attn_paged_prefill` (addresses each request via `qo_indptr` + `kv_indptr` directly against the flat pool). Falls back to `FlashAttnBackend` when dtype unsupported. - **`FlashAttnBackend`**: optional flash-attn dispatch with `flash_attn_with_kvcache` fast path for contiguous cache; falls back to KV gather + `flash_attn_func`. - **`TorchNativeBackend`** (always-available fallback): writes K/V to cache, gathers via `req_to_token` indirect indexing, calls `F.scaled_dot_product_attention`. - Default priority: cuda > flash > torch. Set `ASTR_BACKEND=cuda|torch_native|flash` to override. -Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/rotary_backend.py`, which auto-dispatches to the fused CUDA kernel (`rotary_emb.cu`) during inference or torch complex multiply during training (for autograd compatibility). Both attention backends share the same rotary dispatch. +Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/backend/rotary.py`, which auto-dispatches to the fused CUDA kernel (`rotary_emb.cu`) during inference or torch complex multiply during training (for autograd compatibility). Both attention backends share the same rotary dispatch. Backend selection is thread-safe via `contextvars`, mirroring `torch.nn.attention.sdpa_kernel`: diff --git a/docs/guides/inference.md b/docs/guides/inference.md index b996417..7544755 100644 --- a/docs/guides/inference.md +++ b/docs/guides/inference.md @@ -90,7 +90,7 @@ Fallback: when `CudaBackend` cannot handle an input (wrong dtype or head_dim), ` ### Rotary Embedding Backend -Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/rotary_backend.py`, which auto-dispatches: +Rotary embedding is applied via `apply_rotary_emb` in `astrai/extension/backend/rotary.py`, which auto-dispatches: - **CUDA kernel** (`rotary_emb.cu`): fused cos/sin lookup + rotation in a single kernel, used when the kernel is available, the input is bf16 on CUDA, and `torch.is_grad_enabled()` is `False` (inference mode) - **Torch fallback**: complex multiply path (`torch.view_as_complex` → `torch.complex` multiply → `torch.view_as_real`), used during training (supports autograd backward) or when the CUDA kernel is not available diff --git a/tests/extension/test_backend.py b/tests/extension/test_backend.py index cdbcb5f..56c7e92 100644 --- a/tests/extension/test_backend.py +++ b/tests/extension/test_backend.py @@ -10,6 +10,8 @@ from astrai.extension import ( ATTN_BACKEND, AttentionBackendFactory, CudaBackend, + FlashAttnBackend, + TorchNativeBackend, attn_backend, get_backend, ) @@ -17,13 +19,6 @@ from astrai.extension import ( def test_default_backend_resolves_to_available(): """Default backend is the first available in cuda > flash > torch order.""" - from astrai.extension.attention_backend import ( - CudaBackend, - FlashAttnBackend, - TorchNativeBackend, - _resolve_default_backend, - ) - backend = get_backend() assert isinstance(backend, (CudaBackend, FlashAttnBackend, TorchNativeBackend)) diff --git a/tests/extension/test_kernel_mask.py b/tests/extension/test_kernel_mask.py index 22d9195..ab1ed76 100644 --- a/tests/extension/test_kernel_mask.py +++ b/tests/extension/test_kernel_mask.py @@ -2,14 +2,13 @@ import torch +from astrai.extension.ops.attention import attn_prefill from tests.extension.conftest import D, skip_no_kernel @skip_no_kernel def test_kernel_accepts_2d_mask(): """Kernel should accept 2D mask [batch, kv_len].""" - from astrai.extension.attention_ops import attn_prefill - batch, q_len, n_heads, n_kv_heads = 1, 8, 4, 1 kv_len = 8 q = torch.randn(batch, q_len, n_heads, D, device="cuda", dtype=torch.bfloat16) @@ -25,8 +24,6 @@ def test_kernel_accepts_2d_mask(): @skip_no_kernel def test_kernel_accepts_3d_mask(): """Kernel should accept 3D mask [batch, q_len, kv_len].""" - from astrai.extension.attention_ops import attn_prefill - batch, q_len, n_heads, n_kv_heads = 1, 8, 4, 1 kv_len = 8 q = torch.randn(batch, q_len, n_heads, D, device="cuda", dtype=torch.bfloat16) @@ -41,8 +38,6 @@ def test_kernel_accepts_3d_mask(): @skip_no_kernel def test_kernel_accepts_4d_mask(): """Kernel should accept 4D mask [batch, n_heads, q_len, kv_len].""" - from astrai.extension.attention_ops import attn_prefill - batch, q_len, n_heads, n_kv_heads = 1, 8, 4, 1 kv_len = 8 q = torch.randn(batch, q_len, n_heads, D, device="cuda", dtype=torch.bfloat16) @@ -58,8 +53,6 @@ def test_kernel_accepts_4d_mask(): @skip_no_kernel def test_4d_mask_matches_no_mask_when_all_true(): """A 4D all-True mask should produce the same output as no mask.""" - from astrai.extension.attention_ops import attn_prefill - batch, q_len, n_heads, n_kv_heads = 1, 8, 4, 1 kv_len = 8 q = torch.randn(batch, q_len, n_heads, D, device="cuda", dtype=torch.bfloat16)