16 Commits
Author SHA1 Message Date
ViperEkura 02469887f5 refactor: simplify inference engine and backend dispatch
- merge _generate_streaming/_generate_non_streaming into single _generate() with stream flag
- delete dead GenerationRequest class and generate_with_request method
- inline _next_token helper into generate_async
- replace flash-attn double-checked locking with functools.lru_cache
- extract _write_and_gather_kv helper shared by TorchNative/FlashAttn backends
- inline _kv_cache_is_contiguous into its sole call site in FlashAttnBackend
- change default backend priority from flash>cuda>torch to cuda>flash>torch
- add ASTR_BACKEND env var to override default backend at resolve time
- add supports_graph() static method to AttentionBackend ABC, override in CudaBackend
- replace isinstance(get_backend(), CudaBackend) with get_backend().supports_graph() in executor
- add torch.cuda.is_available() guard to CudaBackend.supports()
2026-08-07 22:28:48 +08:00
ViperEkura 05739629fc feat: add timed() context manager and backend supports()
- Each backend exposes static supports(**kwargs) for capability query
- CudaBackend.supports checks head_dim + kernel availability
- FlashAttnBackend/TorchNativeBackend always return True
- timed() context manager gated by ASTRAI_TIMED=1 env var, logs via logger.info
- Wraps warmup prefill/decode, execute_prefill, and execute_decode
2026-08-07 20:51:30 +08:00
ViperEkura e0f7fa8e13 feat: enable CUDA graph by default with init-time warmup
- Pre-allocate decode_out in InferenceWorkspace so attn_paged_decode does not call torch::empty inside graph capture
- Run live forward before graph capture to get valid output (graph pool memory is zeroed after capture block exits)
- Greedy generation with graph replay is bit-exact across all batch sizes
- _warmup_cuda_graphs pre-captures graphs at init for power-of-two batch sizes
- Graph enabled only when CudaBackend + supported head_dim + warmup succeeds
- Decode speedups vs no-graph: B=1 2.09x, B=4 1.80x, B=8 1.94x, B=16 1.76x
2026-08-07 20:04:10 +08:00
ViperEkura af25833fab fix: add out_buf to attn_paged_decode for CUDA graph capture compatibility
- Pre-allocate decode_out in InferenceWorkspace so attn_paged_decode does not call torch::empty inside graph capture
- Wire decode_out through KVCache, PagePool.bind_tasks, and CudaBackend.fwd_decode
- Run live forward before graph capture to get valid output (graph pool memory is zeroed after capture block exits)
- Greedy generation with graph replay is bit-exact across all batch sizes
- Decode speedups vs no-graph: B=1 2.09x, B=4 1.80x, B=8 1.94x, B=16 1.76x
2026-08-07 19:45:59 +08:00
ViperEkura 6572be4f98 fix: prevent signal handler test from racing with training completion
- Set n_epoch=99999 so training runs until parent delivers signal instead of finishing too fast on CPU
- Drop ready-file deadline from 30s to 10s
2026-08-07 18:32:08 +08:00
ViperEkura 81788faef4 perf: use flash_attn_with_kvcache for contiguous cache decode
- Decode with contiguous cache uses flash_attn_with_kvcache instead of materializing full KV via gather + flash_attn_func
- _backend_supports allows FlashAttnBackend for decode (q_len==1) even with explicit mask
- Decode speedups vs TorchNative (B=1,4,8,16 mean): cuda 1.55x, flash 1.40x, torch_native 1.00x
- Read K/V directly from flat pool via cache_batch_idx + cache_seqlens, zero-copy view reshape
2026-08-07 18:21:02 +08:00
ViperEkura 0e7fe57d96 fix: use max_context_len for stable num_splits in paged decode
- PagedKV::host_kv_len now returns max_context_len instead of max_seq_len
- Eliminates grid-z instability for CUDA graph capture/replay
- Restore skip_no_kernel re-export accidentally removed by ruff --fix
2026-08-07 14:42:53 +08:00
ViperEkura 55ee258e95 style: fix ruff lint warnings
- Remove unused local variable b in attention_backend.py
- Remove unused variable rank0_sd in test_broadcast_state_dict.py
- Remove unused imports across test files
2026-08-07 14:17:48 +08:00
ViperEkura ef1bb6f401 refactor: unify greedy check with _is_greedy helper
- Replace batch-scattered temperature==0 checks with (temperature == 0).all()

- Reuse _is_greedy in standalone sample() function
2026-08-07 14:14:15 +08:00
ViperEkura 6f49738991 feat: auto-select best available attention backend
- Default backend resolves to highest-priority available: flash -> cuda -> torch
- attention() falls back per-call for training/fp32/unsupported head_dim
- Re-apply index_copy_ for CUDA KV cache writes (index_put_ race mitigation)
2026-08-07 13:48:59 +08:00
ViperEkura a59ae8f32e fix: use c10::optional for o_part_buf/ml_part_buf decode kernel params 2026-08-06 20:50:48 +08:00
ViperEkura 6054b8dbd4 feat: add CUDA-graph capture for decode forward
- New CudaGraphContext class: warmup -> capture -> replay lifecycle
- One graph per batch_size key, all inputs at fixed workspace addresses
- Added position_ids buffer to InferenceWorkspace (required for graph capture)
- Graph only activates when CUDA backend is the current backend
- Default off (opt-in) due to slight numerical divergence in graph replay
- Sampling stays outside the graph (torch.multinomial uses mutable RNG)
- Resolved circular import: KVCache -> TYPE_CHECKING in attention_backend.py
2026-08-06 19:57:12 +08:00
ViperEkura 6f67ba8942 perf: move decode split partials to InferenceWorkspace
- Replace per-.cu-file static cached tensors with workspace-managed pre-allocated buffers

- InferenceWorkspace now owns decode_o_part / decode_ml_part (mirrors FlashInfer's workspace pattern)

- KVCache carries the buffers through the backend -> C++ kernel chain

- C++ kernels accept optional pre-allocated buffers; fallback to alloc_split_partials for backward compat

- Pre-allocates once at Executor init, zero allocation in the decode hot loop

- Prerequisite for CUDA-graph capture (all kernel addresses are stable)
2026-08-06 19:12:09 +08:00
ViperEkura d0c5debbab perf: preload V in decode split-kv shared mem and cache partial tensors
- Preload V into shared memory alongside K to eliminate per-element KV address lookups in the inner softmax/accum loop (doubles smem)
- Cache split-KV partial tensors (o_part, ml_part) with static tensors instead of per-call allocation in both decode and paged-decode paths
- Force is_causal=True in CUDA decode backend (decode is always causal)
2026-08-06 18:27:00 +08:00
ViperEkura 4f2e03880b fix : repair and extend throughput benchmark
- adapt bind_tasks to workspace API and reuse a stable workspace
- drop required checkpoint, randomize default 1B GQA preset
- add config override flag for arbitrary model architectures
2026-08-06 12:51:30 +08:00
ViperEkura 5c180cfa90 fix : handle zero-token batch generation
- return empty results without running inference for non-positive limits
- keep scheduler batch outputs aligned with requested max_tokens
- add engine and scheduler regression coverage
2026-08-06 12:31:09 +08:00
27 changed files with 942 additions and 418 deletions
-2
View File
@@ -21,7 +21,6 @@ from astrai.dataset import (
)
from astrai.factory import BaseFactory
from astrai.inference import (
GenerationRequest,
InferenceEngine,
ProtocolHandler,
SamplingPipeline,
@@ -98,7 +97,6 @@ __all__ = [
"EmbeddingEncoder",
"EncoderConfig",
"ExecutorFactory",
"GenerationRequest",
"InferenceEngine",
"LoRAConfig",
"Pipeline",
+219 -118
View File
@@ -22,7 +22,8 @@ Usage — mirroring ``torch.nn.attention.sdpa_kernel``:
Thread-safe via ``contextvars`` — each scheduler thread gets its own
active backend. ``get_backend()`` returns the active one, falling back
to a process-wide ``TorchNativeBackend`` singleton.
to a process-wide default (cuda > flash > torch, overridable via
``ASTR_BACKEND``).
Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
(blhd). The backend returns ``[batch, seq_len, n_heads * head_dim]``.
@@ -30,11 +31,12 @@ Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
import contextvars
import enum
import functools
import importlib
import threading
import os
from abc import ABC, abstractmethod
from contextlib import contextmanager
from typing import Optional, Union
from typing import TYPE_CHECKING, Optional, Union
import torch
import torch.nn.functional as F
@@ -44,67 +46,25 @@ from astrai.extension.attention_ops import (
attn_paged_decode,
attn_paged_prefill,
)
from astrai.extension.loader import is_available
from astrai.factory import BaseFactory
from astrai.inference.core.cache import KVCache
if TYPE_CHECKING:
from astrai.inference.core.cache import KVCache
_current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.ContextVar(
"attn_backend"
)
_lock = threading.Lock()
_flash_available: Optional[bool] = None
@functools.lru_cache(maxsize=1)
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()
@@ -114,8 +74,6 @@ def _flash_attn_check() -> bool:
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
@@ -126,6 +84,14 @@ def _flash_attn_check() -> 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``."""
@@ -134,15 +100,80 @@ class ATTN_BACKEND(enum.Enum):
FLASH = "flash"
_default_backend: Optional["AttentionBackend"] = None
def _priority_backends() -> list["AttentionBackend"]:
"""Available backends in priority order: cuda -> flash -> torch."""
backends: list[AttentionBackend] = []
if is_available("attn_paged_decode") and is_available("attn_paged_prefill"):
backends.append(CudaBackend())
if flash_attn_available():
backends.append(FlashAttnBackend())
backends.append(TorchNativeBackend())
return backends
def _backend_supports(
backend: "AttentionBackend",
q: Tensor,
kv_cache: Optional["KVCache"],
attn_mask: Optional[Tensor],
is_causal: bool,
) -> bool:
"""Whether ``backend`` can run this attention call.
The CUDA kernels are bf16-only, support head_dim in 32/64/128/256, and
need a KV cache (decode/prefill); everything else falls back to torch.
"""
if isinstance(backend, CudaBackend):
return (
kv_cache is not None
and q.dtype == torch.bfloat16
and q.size(-1) in (32, 64, 128, 256)
)
if isinstance(backend, FlashAttnBackend):
if not flash_attn_available():
return False
if q.size(1) == 1 and kv_cache is not None:
return True
return not (attn_mask is not None and not is_causal)
return True
def _resolve_default_backend() -> "AttentionBackend":
"""Pick the highest-priority available backend (cuda -> flash -> torch).
Set ``ASTR_BACKEND`` to override: ``ASTR_BACKEND=cuda``, ``torch_native``,
or ``flash``. The value is the registered name (same as the
``ATTN_BACKEND`` enum value).
Resolved lazily on first ``get_backend()`` and cached. Per-call
capability fallback happens in ``attention()``, so the default is
safe for training and fp32 models.
"""
forced = os.environ.get("ASTR_BACKEND", "").strip().lower()
if forced:
try:
return AttentionBackendFactory.create(forced)
except (ValueError, RuntimeError):
pass
return _priority_backends()[0]
def get_backend() -> "AttentionBackend":
"""Return the active backend for the current thread/context.
Falls back to a ``TorchNativeBackend`` singleton when no backend
has been activated via ``with``.
Falls back to the highest-priority available backend (cuda -> flash ->
torch_native) when no backend has been activated via ``with``. Set
``ASTR_BACKEND`` to override the default.
"""
try:
return _current_backend.get()
except LookupError:
global _default_backend
if _default_backend is None:
_default_backend = _resolve_default_backend()
return _default_backend
@@ -195,11 +226,33 @@ def repeat_kv(x: Tensor, n_rep: int) -> Tensor:
)
def _write_and_gather_kv(
kv_cache: "KVCache",
k: Tensor,
v: Tensor,
layer_id: int,
q: Tensor,
attn_mask: Optional[Tensor],
) -> tuple[Tensor, Tensor]:
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))
return kv_cache.k_buffer[layer_id, indices], kv_cache.v_buffer[layer_id, indices]
def attention(
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache] = None,
kv_cache: Optional["KVCache"] = None,
layer_id: int = 0,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
@@ -223,6 +276,16 @@ def attention(
[batch, q_len, n_heads * head_dim]
"""
backend = get_backend()
if not _backend_supports(backend, q, kv_cache, attn_mask, is_causal):
# The active backend cannot run this call (e.g. CUDA on a training /
# fp32 / unsupported-head_dim input) — fall back to the highest-
# priority backend that can, ending at torch SDPA.
for candidate in _priority_backends():
if isinstance(candidate, type(backend)):
continue
if _backend_supports(candidate, q, kv_cache, attn_mask, is_causal):
backend = candidate
break
return backend.forward(q, k, v, kv_cache, layer_id, attn_mask, is_causal)
@@ -255,7 +318,7 @@ class AttentionBackend(ABC):
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
@@ -284,7 +347,7 @@ class AttentionBackend(ABC):
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
@@ -297,13 +360,24 @@ class AttentionBackend(ABC):
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
) -> Tensor:
"""Multi-token prefill or training forward."""
@staticmethod
def supports_graph() -> bool:
"""Return True if this backend supports CUDA-graph capture.
Override in subclasses that can run under ``torch.cuda.graph``.
Called on the *active* backend instance (or its class) — a cheap
boolean check with no side-effects.
"""
return False
class AttentionBackendFactory(BaseFactory[AttentionBackend]):
"""Factory for registered attention backends."""
@@ -321,12 +395,16 @@ class TorchNativeBackend(AttentionBackend):
runs SDPA directly on the projected q/k/v.
"""
@staticmethod
def supports(**kwargs) -> bool:
return True
def fwd_decode(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
@@ -338,7 +416,7 @@ class TorchNativeBackend(AttentionBackend):
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
@@ -350,30 +428,13 @@ class TorchNativeBackend(AttentionBackend):
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
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]
# Zero out padding positions so gather never touches invalid slots.
# Decode: attn_mask[:,0,0] is exactly the per-position validity
# mask ([B, max_len], True=keep). Prefill: fall back to seq_lens.
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]
k, v = _write_and_gather_kv(kv_cache, k, v, layer_id, q, attn_mask)
n_rep = q.size(2) // k.size(2)
if n_rep > 1:
@@ -391,9 +452,6 @@ class TorchNativeBackend(AttentionBackend):
return out
_default_backend = TorchNativeBackend()
@AttentionBackendFactory.register(ATTN_BACKEND.CUDA.value)
class CudaBackend(AttentionBackend):
"""CUDA kernel backend with direct KV cache access.
@@ -405,18 +463,33 @@ class CudaBackend(AttentionBackend):
``attn_paged_prefill`` with ragged-batch support via qo_indptr +
kv_indptr.
``kv_cache is None`` (training) is not handled — use
``TorchNativeBackend`` for training.
``kv_cache is None`` (training) raises — the per-call fallback to
torch SDPA for training / fp32 / unsupported head_dim happens in the
``attention()`` entry point.
Raises ``RuntimeError`` if the required kernel is not available.
"""
@staticmethod
def supports(**kwargs) -> bool:
head_dim = kwargs.get("head_dim", -1)
return (
torch.cuda.is_available()
and head_dim in (32, 64, 128, 256)
and is_available("attn_paged_decode")
and is_available("attn_paged_prefill")
)
@staticmethod
def supports_graph() -> bool:
return True
def fwd_decode(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
@@ -424,10 +497,10 @@ class CudaBackend(AttentionBackend):
if kv_cache is None:
raise RuntimeError("CudaBackend does not support training (kv_cache=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
loc = kv_cache.out_cache_loc[:, 0]
kv_cache.k_buffer[layer_id].index_copy_(0, loc, k[:, 0])
kv_cache.v_buffer[layer_id].index_copy_(0, loc, v[:, 0])
b = q.size(0)
q_3d = q.squeeze(1)
kv_indptr = kv_cache.kv_indptr
@@ -440,8 +513,10 @@ class CudaBackend(AttentionBackend):
kv_cache.req_pool_indices,
kv_indptr,
kv_cache.max_len,
mask=attn_mask,
is_causal=is_causal,
is_causal=True,
o_part_buf=kv_cache.decode_o_part,
ml_part_buf=kv_cache.decode_ml_part,
out_buf=kv_cache.decode_out,
)
return out.unsqueeze(1).flatten(2)
@@ -450,7 +525,7 @@ class CudaBackend(AttentionBackend):
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
@@ -458,8 +533,13 @@ class CudaBackend(AttentionBackend):
if kv_cache is None:
raise RuntimeError("CudaBackend does not support training (kv_cache=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
loc = kv_cache.out_cache_loc.reshape(-1)
kv_cache.k_buffer[layer_id].index_copy_(
0, loc, k.reshape(-1, k.size(2), k.size(3))
)
kv_cache.v_buffer[layer_id].index_copy_(
0, loc, v.reshape(-1, v.size(2), v.size(3))
)
b = q.size(0)
q_len = q.size(1)
@@ -486,23 +566,26 @@ class CudaBackend(AttentionBackend):
@AttentionBackendFactory.register(ATTN_BACKEND.FLASH.value)
class FlashAttnBackend(AttentionBackend):
"""FlashAttention (FA2/FA3) backend via the optional ``flash-attn`` package.
"""FlashAttention 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.
Decode (q_len=1, contiguous cache): uses ``flash_attn_with_kvcache``,
which reads K/V directly from the flat pool via cache_batch_idx +
cache_seqlens — no materialized KV gather.
For a torch fallback, select ``TorchNativeBackend`` instead.
Prefill / non-contiguous decode: falls back to KV gather +
``flash_attn_func``.
"""
@staticmethod
def supports(**kwargs) -> bool:
return flash_attn_available()
def fwd_decode(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
@@ -514,7 +597,7 @@ class FlashAttnBackend(AttentionBackend):
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
kv_cache: Optional["KVCache"],
layer_id: int,
attn_mask: Optional[Tensor] = None,
is_causal: bool = False,
@@ -526,27 +609,17 @@ class FlashAttnBackend(AttentionBackend):
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: Optional[KVCache],
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]
if q.size(1) == 1 and kv_cache.k_buffer.size(
1
) == kv_cache.req_to_token.size(0) * kv_cache.req_to_token.size(1):
return self._decode_with_kvcache(q, k, v, kv_cache, layer_id)
k, v = _write_and_gather_kv(kv_cache, k, v, layer_id, q, attn_mask)
n_rep = q.size(2) // k.size(2)
if n_rep > 1:
@@ -568,3 +641,31 @@ class FlashAttnBackend(AttentionBackend):
q.contiguous(), k.contiguous(), v.contiguous(), causal=is_causal
)
return out.contiguous().flatten(2)
def _decode_with_kvcache(
self,
q: Tensor,
k: Tensor,
v: Tensor,
kv_cache: "KVCache",
layer_id: int,
) -> Tensor:
max_batch = kv_cache.req_to_token.size(0)
max_seq = kv_cache.req_to_token.size(1)
n_kv = k.size(2)
k_cache = kv_cache.k_buffer[layer_id].view(max_batch, max_seq, n_kv, k.size(3))
v_cache = kv_cache.v_buffer[layer_id].view(max_batch, max_seq, n_kv, v.size(3))
fa = _get_flash_attn()
out = fa.flash_attn_with_kvcache(
q=q,
k_cache=k_cache,
v_cache=v_cache,
k=k,
v=v,
cache_seqlens=(kv_cache.seq_lens - 1).to(torch.int32),
cache_batch_idx=kv_cache.req_pool_indices.to(torch.int32),
causal=True,
)
return out.flatten(2)
+9
View File
@@ -100,6 +100,9 @@ def attn_paged_decode(
max_seq_len: int,
mask: Optional[torch.Tensor] = None,
is_causal: bool = False,
o_part_buf: Optional[torch.Tensor] = None,
ml_part_buf: Optional[torch.Tensor] = None,
out_buf: Optional[torch.Tensor] = None,
) -> torch.Tensor:
"""SGLang-style paged decode (q_len == 1, flat KV pool).
@@ -117,6 +120,9 @@ def attn_paged_decode(
max_seq_len: max per-request seq_len (Python int, for split computation)
mask: 2D [batch, max_seq_len] (bool, True=keep) or None
is_causal: apply causal mask
o_part_buf: pre-allocated split-KV o partial buffer (workflow bypass)
ml_part_buf: pre-allocated split-KV m/l buffer (workflow bypass)
out_buf: pre-allocated output buffer [batch, n_heads, head_dim] (graph-safe)
Returns:
[batch, n_heads, head_dim] (bf16, 3D)
@@ -133,6 +139,9 @@ def attn_paged_decode(
max_seq_len,
mask=mask,
causal_offset=causal_offset,
o_part_buf=o_part_buf,
ml_part_buf=ml_part_buf,
out_buf=out_buf,
)
+2 -3
View File
@@ -5,7 +5,7 @@ Layers:
- api/: HTTP orchestration (ProtocolHandler, server)
- protocols/: Response builders (OpenAI, Anthropic)
- transport/: SSE transport utilities
- engine.py: Facade (InferenceEngine), Value Object (GenerationRequest)
- engine.py: Facade (InferenceEngine)
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy, FrequencyPenaltyStrategy)
"""
@@ -42,7 +42,7 @@ from astrai.inference.core import (
TaskStatus,
page_hash,
)
from astrai.inference.engine import GenerationRequest, InferenceEngine
from astrai.inference.engine import InferenceEngine
from astrai.inference.sample import (
BaseSamplingStrategy,
FrequencyPenaltyStrategy,
@@ -55,7 +55,6 @@ from astrai.inference.sample import (
__all__ = [
"InferenceEngine",
"GenerationRequest",
"InferenceScheduler",
"Executor",
"STOP",
+15
View File
@@ -260,6 +260,10 @@ class KVCache:
max_len: max(seq_lens) as Python int — avoids GPU sync in decode
kv_indptr: [batch+1] int32 — prefix sum of seq_lens, precomputed once
per step so the attention backend avoids rebuilding it per layer.
qo_indptr: [batch+1] int32 — prefill qo prefix-sum (None in decode)
decode_o_part: split-KV o partial workspace (mirrors FlashInfer)
decode_ml_part: split-KV m/l partial workspace (mirrors FlashInfer)
decode_out: pre-allocated decode output buffer (graph-safe)
"""
k_buffer: Tensor
@@ -271,6 +275,9 @@ class KVCache:
max_len: int = 0
kv_indptr: Optional[Tensor] = None
qo_indptr: Optional[Tensor] = None
decode_o_part: Optional[Tensor] = None
decode_ml_part: Optional[Tensor] = None
decode_out: Optional[Tensor] = None
class PagePool:
@@ -565,12 +572,17 @@ class PagePool:
torch.arange(b + 1, dtype=torch.int32, device=device) * q_len
)
qo_indptr = workspace.qo_indptr[: b + 1]
decode_o_part, decode_ml_part = None, None
decode_out = None
else:
write_pos = seq_lens_t - 1
loc = self._req_pool.req_to_token[req_pool_indices, write_pos].unsqueeze(-1)
ocl_buf[:b].copy_(loc)
out_cache_loc = ocl_buf[:b]
qo_indptr = None
decode_o_part = getattr(workspace, "decode_o_part", None)
decode_ml_part = getattr(workspace, "decode_ml_part", None)
decode_out = getattr(workspace, "decode_out", None)
return KVCache(
k_buffer=self._storage.k_buffer,
@@ -582,6 +594,9 @@ class PagePool:
max_len=max(seq_lens),
kv_indptr=kv_indptr,
qo_indptr=qo_indptr,
decode_o_part=decode_o_part,
decode_ml_part=decode_ml_part,
decode_out=decode_out,
)
# ---- internals ----
+177 -22
View File
@@ -1,11 +1,21 @@
import logging
import os
import time
from contextlib import contextmanager
from dataclasses import dataclass
from typing import List, Optional
import torch
from torch import Tensor
from astrai.extension.attention_backend import (
ATTN_BACKEND,
CudaBackend,
attn_backend,
get_backend,
)
from astrai.inference.core.cache import PagePool
from astrai.inference.core.graph import CudaGraphContext
from astrai.inference.core.task import Task
from astrai.inference.core.workspace import InferenceWorkspace
from astrai.inference.sample import sample
@@ -13,6 +23,19 @@ from astrai.model.automodel import AutoModel
from astrai.tokenize.tokenizer import AutoTokenizer
logger = logging.getLogger(__name__)
_TIMED = os.environ.get("ASTRAI_TIMED", "") == "1"
@contextmanager
def timed(label: str, log: Optional[logging.Logger] = None):
"""Wall-clock debug timer, enabled via ``ASTRAI_TIMED=1``."""
if not _TIMED:
yield
return
tic = time.perf_counter()
yield
elapsed_ms = (time.perf_counter() - tic) * 1000
(log or logger).info("%s %.1fms", label, elapsed_ms)
@dataclass
@@ -51,6 +74,80 @@ def _build_sampling_batch_info(tasks: List[Task], device) -> SamplingBatchInfo:
)
def _warmup_cuda_graphs(
model: AutoModel,
pool: PagePool,
ws: InferenceWorkspace,
gctx: "CudaGraphContext",
max_batch_size: int,
prompt_len: int = 1,
device: Optional[str] = None,
):
batch_sizes = [1]
n = 2
while n <= max_batch_size:
batch_sizes.append(n)
n *= 2
if max_batch_size not in batch_sizes:
batch_sizes.append(max_batch_size)
dev = device or next(model.parameters()).device
for b in batch_sizes:
task_ids = [f"_gr_{b}_{i}" for i in range(b)]
prompt_tokens = [list(range(prompt_len)) for _ in range(b)]
alloc_ok = True
for tid, pt in zip(task_ids, prompt_tokens):
if not pool.task_alloc(tid, pt):
alloc_ok = False
break
if not alloc_ok:
for tid in task_ids:
pool.task_free(tid)
continue
with (
torch.inference_mode(),
attn_backend(ATTN_BACKEND.CUDA),
timed(f"warmup prefill b={b}", logger),
):
kv_cache = pool.bind_tasks(task_ids, ws, start_pos=0)
ids_in = torch.tensor(prompt_tokens, dtype=torch.long, device=dev)
pos_in = torch.arange(prompt_len, device=dev).unsqueeze(0).expand(b, -1)
model(
ids_in,
input_mask=pos_in.unsqueeze(-1) >= torch.arange(prompt_len, device=dev),
kv_cache=kv_cache,
position_ids=pos_in,
)
with (
torch.inference_mode(),
attn_backend(ATTN_BACKEND.CUDA),
timed(f"warmup decode b={b}", logger),
):
for step in range(2):
seq_pos = prompt_len + step
ws.position_ids[:b] = seq_pos
for tid in task_ids:
pool.task_extend(tid, seq_pos)
kv = pool.bind_tasks(task_ids, ws)
input_mask = ws.decode_mask(ws.position_ids[:b], ws.max_seq_len)
ids_buf = ws.fill_input_ids([step] * b)
gctx.forward(
model,
key=(b,),
input_ids=ids_buf.unsqueeze(1),
input_mask=input_mask,
kv_cache=kv,
position_ids=ws.position_ids[:b].unsqueeze(1),
)
for tid in task_ids:
pool.task_free(tid)
torch.cuda.synchronize()
class Executor:
"""Model forward passes for prefill and decode phases."""
@@ -78,13 +175,40 @@ class Executor:
# (input_ids, decode mask, KV bind metadata). Eagerly sized at init
# so the workspace is CUDA-graph-capture friendly — no allocation
# during capture.
config = model.config
max_q_heads = config.num_attention_heads
head_dim = config.hidden_size // config.num_attention_heads
self._head_dim = head_dim
self._graph_supported = CudaBackend.supports(head_dim=head_dim)
self._workspace = InferenceWorkspace(
max_batch_size=kv_cache.max_batch_size,
max_seq_len=kv_cache.max_seq_len,
max_q_heads=max_q_heads,
head_dim=head_dim,
device=self.device,
dtype=self.dtype,
)
# CUDA-graph capture: one graph per (batch_size,) key.
# Enabled at init-time via _warmup_cuda_graphs for CudaBackend
# on supported head_dims; left disabled otherwise.
self._graph_ctx = CudaGraphContext()
self._try_enable_cuda_graph()
def _try_enable_cuda_graph(self):
if not self._graph_supported:
return
self._graph_ctx.set_enabled(True)
_warmup_cuda_graphs(
self.model,
self.kv_cache,
self._workspace,
self._graph_ctx,
max_batch_size=self.kv_cache.max_batch_size,
device=self.device,
)
def _sample_logits(
self,
logits: Tensor,
@@ -164,7 +288,10 @@ class Executor:
prompt_len, device=self.device
)
with torch.inference_mode():
with (
torch.inference_mode(),
timed(f"execute_prefill b={batch_sz} prompt_len={prompt_len}", logger),
):
outputs = self.model(
input_ids,
input_mask=input_mask,
@@ -199,43 +326,71 @@ class Executor:
if not tasks:
return []
input_ids = self._workspace.fill_input_ids(
b = len(tasks)
ws = self._workspace
# ---- pre-replay: update input buffers in-place ----
input_ids = ws.fill_input_ids(
[t.output_ids[-1] if t.output_ids else t.prompt_ids[-1] for t in tasks]
).unsqueeze(1)
)
task_ids = [t.task_id for t in tasks]
cur_positions = [t.next_pos for t in tasks]
sig = tuple(task_ids)
cur_positions = [t.next_pos for t in tasks]
cached = self._decode_cache
if (
cached is not None
and cached[0] == sig
and cur_positions == [p + 1 for p in cached[1]]
):
_, _, info, position_ids = cached
position_ids += 1
self._decode_cache = (sig, cur_positions, info, position_ids)
info = cached[2]
ws.position_ids[:b] += 1
self._decode_cache = (sig, cur_positions, info)
else:
info = _build_sampling_batch_info(tasks, self.device)
position_ids = torch.tensor(
cur_positions, dtype=torch.long, device=self.device
ws.position_ids[:b].copy_(
torch.tensor(cur_positions, dtype=torch.long, device=self.device)
)
self._decode_cache = (sig, cur_positions, info, position_ids)
self._decode_cache = (sig, cur_positions, info)
total_len = max(t.next_pos for t in tasks) + 1
input_mask = self._workspace.decode_mask(position_ids, total_len)
total_len = max(cur_positions) + 1
input_mask = ws.decode_mask(ws.position_ids[:b], total_len)
with torch.inference_mode():
outputs = self.model(
input_ids,
input_mask=input_mask,
kv_cache=self.kv_cache.bind_tasks(
task_ids,
self._workspace,
),
position_ids=position_ids.unsqueeze(1),
)
kv_cache = self.kv_cache.bind_tasks(task_ids, ws)
# ---- forward (graph replay or live run + capture) ----
use_graph = (
self._graph_ctx.enabled
and self._graph_supported
and get_backend().supports_graph()
)
key = (b,)
if use_graph:
input_mask = ws.decode_mask(ws.position_ids[:b], ws.max_seq_len)
with (
torch.inference_mode(),
timed(f"execute_decode forward b={b}", logger),
):
if use_graph:
outputs = self._graph_ctx.forward(
self.model,
key=key,
input_ids=input_ids.unsqueeze(1),
input_mask=input_mask,
kv_cache=kv_cache,
position_ids=ws.position_ids[:b].unsqueeze(1),
)
else:
outputs = self.model(
input_ids.unsqueeze(1),
input_mask=input_mask,
kv_cache=kv_cache,
position_ids=ws.position_ids[:b].unsqueeze(1),
)
logits = outputs["logits"][:, -1, :]
return self._sample_logits(logits, tasks, return_logprobs, info=info)
+103
View File
@@ -0,0 +1,103 @@
"""CUDA-graph capture for the decode model-forward step.
Mirrors SGLang's cuda-graph manager: one graph per batch size. The graph
pair. The graph captures ``model.forward()`` with workspace-backed inputs
(all at fixed addresses). Before each replay the caller updates the input
buffer content in-place so the graph sees fresh data at the same tensor
addresses.
Only the model forward is captured — sampling runs outside the graph
(via ``torch.multinomial`` which consumes a mutable RNG state).
"""
import torch
from torch import Tensor
class CudaGraphContext:
"""CUDA-graph capture/replay for decode steps.
Parameters:
enabled: When ``False``, ``forward()`` always runs the live model
forward without capture/replay (graphs are cleared). Toggle at
runtime via the ``set_enabled()`` method.
Usage::
gctx = CudaGraphContext()
with torch.inference_mode():
outputs = gctx.forward(
model,
key=(batch_size,),
input_ids=workspace.input_ids[:b].unsqueeze(1),
input_mask=input_mask,
kv_cache=kv_cache,
position_ids=workspace.position_ids[:b].unsqueeze(1),
)
The first call at a given key runs *without* capture (warmup). The
second call captures the graph. Subsequent calls replay the captured
graph. A ``torch.cuda.synchronize()`` before capture drains in-flight
work so the graph trace is clean.
"""
def __init__(self, enabled: bool = False):
self._enabled = enabled
self._graphs: dict[tuple, torch.cuda.CUDAGraph] = {}
self._outputs: dict[tuple, dict[str, Tensor]] = {}
self._warmed: set[tuple] = set()
@property
def enabled(self) -> bool:
return self._enabled
def set_enabled(self, flag: bool):
"""Enable or disable CUDA-graph capture at runtime.
Disabling clears all captured graphs (frees GPU memory) and warmup
state. Re-enabling after disable starts fresh — graphs are
re-captured on the next warmup cycle.
"""
if flag == self._enabled:
return
self._enabled = flag
if not flag:
self._graphs.clear()
self._outputs.clear()
self._warmed.clear()
def forward(self, model, *, key, **kwargs) -> dict[str, Tensor]:
"""Run ``model(**kwargs)`` via graph replay or live forward.
Args:
model: callable, e.g. ``self.model.forward``.
key: ``(batch_size,)`` — the dispatch key (one graph per batch size).
**kwargs: arguments forwarded to ``model``. All tensor arguments
must reside at stable addresses (workspace buffers).
Returns:
The dict produced by ``model(**kwargs)``, e.g.
``{"logits": ..., "h0": ...}``.
"""
if not self._enabled:
self._outputs[key] = model(**kwargs)
return self._outputs[key]
if key in self._graphs:
self._graphs[key].replay()
elif key in self._warmed:
cap_output = model(**kwargs)
torch.cuda.synchronize()
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
self._outputs[key] = model(**kwargs)
self._graphs[key] = graph
self._warmed.discard(key)
return cap_output
else:
self._warmed.add(key)
self._outputs[key] = model(**kwargs)
return self._outputs[key]
def has_graph(self, key: tuple) -> bool:
return key in self._graphs
+3
View File
@@ -288,6 +288,9 @@ class InferenceScheduler:
t_max = seq_cap - len(ids)
else:
t_max = min(t_max, seq_cap - len(ids))
if t_max <= 0:
tasks.append(None)
continue
task = Task(
task_id=f"batch_{uuid.uuid4().hex[:8]}",
prompt_ids=list(ids),
+50 -10
View File
@@ -1,20 +1,16 @@
"""Pre-allocated buffers for the inference decode hot path.
Mirrors SGLang's pre-allocated input buffers (``input_buffers.py``): tensors
are sized once to the server's maximum dimensions and sliced to the live
batch each step, so the per-token decode loop never calls
``torch.empty``/``torch.zeros``/``torch.arange`` for the hot shapes. Fills
go through ``out=`` variants (``torch.ge``) which write into the stable
buffers instead of allocating fresh results.
All buffers are allocated eagerly at init (nothing is lazy), so the
workspace is CUDA-graph-capture friendly: the decode step reads/writes
fixed-address tensors with no allocation during capture.
Mirrors FlashInfer / SGLang's global workspace pattern: all per-step tensors
are allocated eagerly at init (nothing is lazy), so the decode step
reads/writes fixed-address tensors with zero ``torch.empty`` calls during
the hot loop — a prerequisite for CUDA-graph capture.
"""
import torch
from torch import Tensor
_MAX_SPLITS = 32
class InferenceWorkspace:
"""Reusable fixed-shape per-step buffers for decode.
@@ -30,6 +26,11 @@ class InferenceWorkspace:
- KV-cache bind metadata (``req_pool_indices``, ``seq_lens``,
``kv_indptr``, ``inc``, ``out_cache_loc``), written by
``PagePool.bind_tasks`` when the Executor passes this workspace.
- ``decode_o_part`` / ``decode_ml_part``: split-KV partial result buffers
(mirrors FlashInfer's workspace). One global alloc, reused by every
decode step across all layers. Sliced views are passed to the CUDA
attention kernel so its internal ``torch.empty`` hot-path alloc goes
through a stable address (CUDA-graph capturable).
No re-allocation while the server's bounds are respected.
"""
@@ -38,11 +39,15 @@ class InferenceWorkspace:
self,
max_batch_size: int,
max_seq_len: int,
max_q_heads: int,
head_dim: int,
device: torch.device,
dtype: torch.dtype,
):
self.max_batch_size = max_batch_size
self.max_seq_len = max_seq_len
self.max_q_heads = max_q_heads
self.head_dim = head_dim
self.device = device
self.dtype = dtype
@@ -83,6 +88,41 @@ class InferenceWorkspace:
(max_batch_size, 1), dtype=torch.long, device=device
)
# Per-step position IDs (must be at a fixed address for CUDA-graph capture).
self.position_ids = torch.empty(
(max_batch_size,), dtype=torch.long, device=device
)
# Split-KV partial-result buffers for decode (persistent, one global
# alloc per process — mirrors FlashInfer's workspace pattern).
# Shape: [max_batch_size, max_q_heads, _MAX_SPLITS, head_dim] (o_part)
# [max_batch_size, max_q_heads, _MAX_SPLITS, 2] (ml_part)
self.decode_o_part = torch.empty(
(max_batch_size, max_q_heads, _MAX_SPLITS, head_dim),
dtype=torch.float32,
device=device,
)
self.decode_ml_part = torch.empty(
(max_batch_size, max_q_heads, _MAX_SPLITS, 2),
dtype=torch.float32,
device=device,
)
# Decode output buffer (graph-safe pre-alloc). Shape matches the
# decode kernel's output: [batch, q_head, head_dim].
self.decode_out = torch.empty(
(max_batch_size, max_q_heads, head_dim),
dtype=dtype,
device=device,
)
def decode_buffers(self, batch: int, q_heads: int):
"""Return ``(o_part, ml_part)`` view sliced to live dimensions."""
return (
self.decode_o_part[:batch, :q_heads],
self.decode_ml_part[:batch, :q_heads],
)
def fill_input_ids(self, ids: "list[int]") -> Tensor:
"""Write ``ids`` into the device buffer and return ``[B]``.
+42 -155
View File
@@ -64,44 +64,6 @@ class GenerateResult:
return self.results.copy()
class GenerationRequest:
"""Request parameters for text generation."""
def __init__(
self,
messages: List[Dict[str, str]],
top_k: int = 50,
top_p: float = 1.0,
temperature: float = 1.0,
max_tokens: Optional[int] = None,
frequency_penalty: float = 0.0,
rep_window: int = 64,
stream: bool = False,
):
if not (isinstance(top_k, int) and top_k >= 0):
raise ValueError("top_k must be a non-negative integer")
if not (0.0 <= top_p <= 1.0):
raise ValueError("top_p must be a float between 0.0 and 1.0")
if not (isinstance(temperature, (int, float)) and temperature >= 0):
raise ValueError("temperature must be a non-negative number")
if not (
isinstance(frequency_penalty, (int, float))
and -2.0 <= frequency_penalty <= 2.0
):
raise ValueError("frequency_penalty must be between -2.0 and 2.0")
if not (isinstance(rep_window, int) and rep_window > 0):
raise ValueError("rep_window must be a positive integer")
self.messages = messages
self.top_k = top_k
self.top_p = top_p
self.temperature = temperature
self.max_tokens = max_tokens
self.frequency_penalty = frequency_penalty
self.rep_window = rep_window
self.stream = stream
class InferenceEngine:
"""Unified inference engine backed by continuous-batching scheduler."""
@@ -146,28 +108,23 @@ class InferenceEngine:
is_batch = isinstance(prompt, list)
prompts = prompt if is_batch else [prompt]
if stream:
return self._generate_streaming(
prompts,
is_batch,
max_tokens,
temperature,
top_p,
top_k,
frequency_penalty,
rep_window,
)
else:
return self._generate_non_streaming(
prompts,
is_batch,
max_tokens,
temperature,
top_p,
top_k,
frequency_penalty,
rep_window,
)
if max_tokens is not None and max_tokens <= 0:
if stream:
return iter(())
results = [""] * len(prompts)
return results if is_batch else results[0]
return self._generate(
prompts,
is_batch,
stream,
max_tokens,
temperature,
top_p,
top_k,
frequency_penalty,
rep_window,
)
def generate_async(
self,
@@ -179,9 +136,10 @@ class InferenceEngine:
frequency_penalty: float = 0.0,
rep_window: int = 64,
) -> AsyncGenerator[str, None]:
sync_gen = self._generate_streaming(
sync_gen = self._generate(
[prompt],
False,
True,
max_tokens,
temperature,
top_p,
@@ -193,51 +151,30 @@ class InferenceEngine:
async def _agen():
loop = asyncio.get_event_loop()
while True:
token = await loop.run_in_executor(None, self._next_token, sync_gen)
if token is None:
try:
token = await loop.run_in_executor(None, next, sync_gen)
except StopIteration:
break
yield token
return _agen()
@staticmethod
def _next_token(gen: Generator) -> Optional[str]:
try:
return next(gen)
except StopIteration:
return None
def generate_with_request(
self, request: GenerationRequest
) -> Union[Generator[str, None, None], str, List[str]]:
prompt = self.tokenizer.apply_chat_template(request.messages, tokenize=False)
return self.generate(
prompt=prompt,
stream=request.stream,
max_tokens=request.max_tokens,
temperature=request.temperature,
top_p=request.top_p,
top_k=request.top_k,
frequency_penalty=request.frequency_penalty,
rep_window=request.rep_window,
)
def _submit_tasks(
def _generate(
self,
prompts: List[str],
is_batch: bool,
stream: bool,
max_tokens: Optional[int],
temperature: float,
top_p: float,
top_k: int,
frequency_penalty: float,
rep_window: int,
) -> Tuple[GenerateResult, List[str]]:
) -> Union[Generator, str, List[str]]:
n = len(prompts)
result = GenerateResult(count=n)
task_ids = []
for i, p in enumerate(prompts):
cb = self._make_callback(result, i)
task_id = self.scheduler.add_task(
task_ids = [
self.scheduler.add_task(
prompt=p,
max_tokens=max_tokens,
temperature=temperature,
@@ -245,39 +182,23 @@ class InferenceEngine:
top_k=top_k,
frequency_penalty=frequency_penalty,
rep_window=rep_window,
stream_callback=cb,
stream_callback=lambda token, idx=i: result.append(token, idx),
)
task_ids.append(task_id)
return result, task_ids
for i, p in enumerate(prompts)
]
@staticmethod
def _make_callback(result: GenerateResult, idx: int):
def cb(token):
result.append(token, idx)
if not stream:
try:
result.wait_completion()
except TimeoutError:
for tid in task_ids:
self.scheduler.remove_task(tid)
raise
for tid in task_ids:
self.scheduler.remove_task(tid)
res = result.get_results()
return res if is_batch else res[0]
return cb
def _generate_streaming(
self,
prompts: List[str],
is_batch: bool,
max_tokens: Optional[int],
temperature: float,
top_p: float,
top_k: int,
frequency_penalty: float,
rep_window: int,
) -> Generator:
result, task_ids = self._submit_tasks(
prompts,
max_tokens,
temperature,
top_p,
top_k,
frequency_penalty,
rep_window,
)
n = len(prompts)
remaining = n
finished = [False] * n
@@ -301,40 +222,6 @@ class InferenceEngine:
return gen()
def _generate_non_streaming(
self,
prompts: List[str],
is_batch: bool,
max_tokens: Optional[int],
temperature: float,
top_p: float,
top_k: int,
frequency_penalty: float,
rep_window: int,
) -> Union[str, List[str]]:
result, task_ids = self._submit_tasks(
prompts,
max_tokens,
temperature,
top_p,
top_k,
frequency_penalty,
rep_window,
)
try:
result.wait_completion()
except TimeoutError:
for tid in task_ids:
self.scheduler.remove_task(tid)
raise
for tid in task_ids:
self.scheduler.remove_task(tid)
res = result.get_results()
return res if is_batch else res[0]
def get_stats(self) -> Dict[str, Any]:
return self.scheduler.get_stats()
+2 -6
View File
@@ -266,7 +266,7 @@ class SamplingPipeline(BaseSamplingStrategy):
@staticmethod
def _is_greedy(temperature: Union[float, Tensor]) -> bool:
if isinstance(temperature, Tensor):
return temperature.numel() == 1 and temperature.item() == 0
return bool((temperature == 0).all())
return temperature == 0
@torch.inference_mode()
@@ -364,11 +364,7 @@ def sample(
``chosen_logprobs`` has shape ``[batch]``.
"""
greedy = (
(
isinstance(temperature, Tensor)
and temperature.numel() == 1
and temperature.item() == 0
)
bool((temperature == 0).all())
if isinstance(temperature, Tensor)
else temperature == 0
)
+17 -2
View File
@@ -8,7 +8,9 @@ torch::Tensor attn_decode(
c10::optional<torch::Tensor> mask,
int64_t causal_offset,
double scale,
int64_t layout
int64_t layout,
c10::optional<torch::Tensor> o_part_buf,
c10::optional<torch::Tensor> ml_part_buf
) {
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
auto stream = at::cuda::getCurrentCUDAStream();
@@ -22,7 +24,18 @@ torch::Tensor attn_decode(
auto O_view = (layout == BLHD) ? O.transpose(1, 2) : O;
p.o = (bf16*)O_view.data_ptr();
alloc_split_partials(p);
if (o_part_buf.has_value() && ml_part_buf.has_value()
&& o_part_buf->defined() && ml_part_buf->defined()) {
TORCH_CHECK(o_part_buf->scalar_type() == torch::kFloat32, "o_part_buf must be f32");
TORCH_CHECK(ml_part_buf->scalar_type() == torch::kFloat32, "ml_part_buf must be f32");
int64_t o_needed = (int64_t)p.batch * p.q_head * MAX_SPLITS * p.head_dim;
TORCH_CHECK(o_part_buf->numel() >= o_needed,
"o_part_buf too small: need ", o_needed, " got ", o_part_buf->numel());
p.o_part = (float*)o_part_buf->data_ptr();
p.ml_part = (float*)ml_part_buf->data_ptr();
} else {
alloc_split_partials(p);
}
DISPATCH_HEAD_DIM(p.head_dim, dispatch_decode, p, stream);
C10_CUDA_CHECK(cudaGetLastError());
return O;
@@ -37,5 +50,7 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
py::arg("causal_offset") = -1,
py::arg("scale") = 0.0,
py::arg("layout") = (int64_t)BHLD,
py::arg("o_part_buf") = py::none(),
py::arg("ml_part_buf") = py::none(),
"GQA decode (tensor-core head-packing on sm_80+, scalar fallback)");
}
+7 -9
View File
@@ -35,7 +35,9 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
float m = -FLT_MAX, d = 0.0f, acc_reg[8] = {0.0f};
extern __shared__ __align__(16) bf16 k_smem[];
extern __shared__ __align__(16) bf16 smem[];
bf16* k_smem = smem;
bf16* v_smem = smem + DC_CHUNK * p.head_dim;
// Split-KV: each split processes a contiguous subset of chunks
int chunks_total = (seq_len + DC_CHUNK - 1) / DC_CHUNK;
@@ -47,8 +49,8 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
int chunk_start = ci * DC_CHUNK;
int this_chunk = min(DC_CHUNK, seq_len - chunk_start);
// Load K into shared memory (addressing via KV policy; paged guards
// empty slots with zero-fill).
// Load K and V into shared memory (addressing via KV policy;
// paged guards empty slots with zero-fill).
int total = this_chunk * p.head_dim;
for (int i = threadIdx.y * 32 + lane; i < total;
i += blockDim.x * blockDim.y) {
@@ -57,6 +59,7 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
int kc = chunk_start + s;
KVAddr a = KV::kv_addr(p, kctx, kc, d_dim, true);
k_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.k) : (bf16)0.f;
v_smem[i] = a.valid ? *reinterpret_cast<const bf16*>(a.v) : (bf16)0.f;
}
__syncthreads();
@@ -82,13 +85,8 @@ __global__ void attn_decode_split_kv_kernel(AttentionParams<bf16> p) {
float beta = __expf(partial - new_m);
d = d * alpha + beta;
// V read via KV policy; when masked (beta == 0) or the slot is
// empty the term vanishes, so no extra branches are needed.
for (int i = 0; i < hd_per_thread; i++) {
KVAddr a = KV::kv_addr(p, kctx, kv_idx, lane * hd_per_thread + i, true);
float vv = a.valid
? __bfloat162float(*reinterpret_cast<const bf16*>(a.v))
: 0.0f;
float vv = __bfloat162float(v_smem[s * p.head_dim + lane * hd_per_thread + i]);
acc_reg[i] = fmaf(acc_reg[i], alpha, vv * beta);
}
m = new_m;
+1 -1
View File
@@ -157,7 +157,7 @@ struct DecodeLauncherScalar {
int kv_len = KV::host_kv_len(p);
int chunks_total = (kv_len + DC_CHUNK - 1) / DC_CHUNK;
p.num_splits = compute_num_splits(p.batch * p.kv_head, chunks_total);
size_t smem = DC_CHUNK * p.head_dim * sizeof(bf16);
size_t smem = 2 * DC_CHUNK * p.head_dim * sizeof(bf16);
int g = min(group_size, 32); // cap at 32 to respect 1024-thread limit
dim3 grid(p.batch * p.kv_head, 1, p.num_splits);
dim3 block(32, g);
+1 -1
View File
@@ -111,7 +111,7 @@ struct PagedKV {
return p.max_q_len;
}
HOST_DEV_FORCEINLINE int host_kv_len(const AttentionParams<bf16>& p) {
return p.max_seq_len;
return p.max_context_len;
}
// prefill: Q rows start at qo_indptr[batch] (ragged batch base)
+31 -3
View File
@@ -11,7 +11,10 @@ torch::Tensor attn_paged_decode(
int64_t max_seq_len,
c10::optional<torch::Tensor> mask,
int64_t causal_offset,
double scale
double scale,
c10::optional<torch::Tensor> o_part_buf,
c10::optional<torch::Tensor> ml_part_buf,
c10::optional<torch::Tensor> out_buf
) {
const at::cuda::OptionalCUDAGuard device_guard(device_of(q));
auto stream = at::cuda::getCurrentCUDAStream();
@@ -21,10 +24,32 @@ torch::Tensor attn_paged_decode(
req_to_token, req_pool_indices, kv_indptr,
max_seq_len, mask, causal_offset, scale, p);
auto O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options());
torch::Tensor O;
if (out_buf.has_value() && out_buf->defined()) {
TORCH_CHECK(out_buf->dtype() == q.dtype(), "out_buf dtype must match q");
TORCH_CHECK(out_buf->size(0) >= q.size(0), "out_buf batch too small");
TORCH_CHECK(out_buf->size(1) >= q.size(1), "out_buf heads too small");
TORCH_CHECK(out_buf->size(2) >= q.size(2), "out_buf head_dim too small");
O = out_buf.value().slice(0, 0, q.size(0))
.slice(1, 0, q.size(1))
.slice(2, 0, q.size(2));
} else {
O = torch::empty({q.size(0), q.size(1), q.size(2)}, q.options());
}
p.o = (bf16*)O.data_ptr();
alloc_split_partials(p);
if (o_part_buf.has_value() && ml_part_buf.has_value()
&& o_part_buf->defined() && ml_part_buf->defined()) {
TORCH_CHECK(o_part_buf->scalar_type() == torch::kFloat32, "o_part_buf must be f32");
TORCH_CHECK(ml_part_buf->scalar_type() == torch::kFloat32, "ml_part_buf must be f32");
int64_t o_needed = (int64_t)p.batch * p.q_head * MAX_SPLITS * p.head_dim;
TORCH_CHECK(o_part_buf->numel() >= o_needed,
"o_part_buf too small: need ", o_needed, " got ", o_part_buf->numel());
p.o_part = (float*)o_part_buf->data_ptr();
p.ml_part = (float*)ml_part_buf->data_ptr();
} else {
alloc_split_partials(p);
}
DISPATCH_HEAD_DIM(p.head_dim, dispatch_paged_decode, p, stream);
C10_CUDA_CHECK(cudaGetLastError());
return O;
@@ -42,5 +67,8 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
py::arg("mask") = py::none(),
py::arg("causal_offset") = -1,
py::arg("scale") = 0.0,
py::arg("o_part_buf") = py::none(),
py::arg("ml_part_buf") = py::none(),
py::arg("out_buf") = py::none(),
"SGLang-style paged decode: flat KV pool + req_to_token + kv_indptr.");
}
+183 -23
View File
@@ -1,3 +1,4 @@
import json
from pathlib import Path
from typing import Optional, Union
@@ -5,15 +6,30 @@ import click
import torch
from astrai import setup_logging
from astrai.config import AutoRegressiveLMConfig
from astrai.config import BaseModelConfig, ConfigFactory
from astrai.extension import ATTN_BACKEND, AttentionBackendFactory, attn_backend
from astrai.inference.core.cache import PagePool
from astrai.model import AutoModel
from astrai.inference.core.graph import CudaGraphContext
from astrai.inference.core.workspace import InferenceWorkspace
from astrai.model import AutoModel, AutoRegressiveLM
_DTYPES = ["bfloat16", "float16", "float32"]
_CACHES = ["contiguous", "paged"]
_BACKENDS = AttentionBackendFactory.list_registered()
# Default 1B GQA preset matching the project checkpoint architecture.
_DEFAULT_CONFIG = {
"vocab_size": 100000,
"hidden_size": 1536,
"num_hidden_layers": 24,
"intermediate_size": 6912,
"num_attention_heads": 24,
"num_key_value_heads": 4,
"max_position_embeddings": 32768,
"rms_norm_eps": 1e-05,
"tie_word_embeddings": False,
}
class BenchmarkResult:
def __init__(
@@ -37,11 +53,12 @@ class GenerationBenchmark:
def __init__(
self,
model: AutoModel,
config: AutoRegressiveLMConfig,
config: BaseModelConfig,
device: str = "cuda",
dtype: torch.dtype = torch.bfloat16,
cache_type: str = "contiguous",
backend: Union[str, ATTN_BACKEND] = ATTN_BACKEND.CUDA,
cuda_graph: bool = False,
):
self.device = device
self.dtype = dtype
@@ -49,6 +66,7 @@ class GenerationBenchmark:
self.model = model
self.config = config
self.backend = backend
self.cuda_graph = cuda_graph
def _make_pool(self, batch_size: int, max_seq_len: int) -> PagePool:
return PagePool(
@@ -63,7 +81,24 @@ class GenerationBenchmark:
n_tokens=None,
)
def _run_prefill(self, pool: PagePool, batch_size: int, prompt_len: int) -> list:
@staticmethod
def _make_workspace(pool: PagePool, config: BaseModelConfig) -> InferenceWorkspace:
return InferenceWorkspace(
pool.max_batch_size,
pool.max_seq_len,
max_q_heads=config.num_attention_heads,
head_dim=config.hidden_size // config.num_attention_heads,
device=pool.device,
dtype=pool.dtype,
)
def _run_prefill(
self,
pool: PagePool,
batch_size: int,
prompt_len: int,
workspace: InferenceWorkspace,
) -> list:
input_ids = torch.randint(
0, self.config.vocab_size, (batch_size, prompt_len), device=self.device
)
@@ -80,9 +115,7 @@ class GenerationBenchmark:
for tid in task_ids:
pool.task_alloc(tid, list(range(prompt_len)))
kv_cache = pool.bind_tasks(
task_ids, [prompt_len] * batch_size, self.device, start_pos=0
)
kv_cache = pool.bind_tasks(task_ids, workspace, self.device, start_pos=0)
with torch.inference_mode(), attn_backend(self.backend):
self.model(
input_ids,
@@ -93,7 +126,13 @@ class GenerationBenchmark:
torch.cuda.synchronize()
return task_ids
def _run_decode_step(self, pool: PagePool, task_ids: list, seq_len: int):
def _run_decode_step(
self,
pool: PagePool,
task_ids: list,
seq_len: int,
workspace: InferenceWorkspace,
):
batch_size = len(task_ids)
input_ids = torch.randint(
0, self.config.vocab_size, (batch_size, 1), device=self.device
@@ -102,10 +141,12 @@ class GenerationBenchmark:
[[seq_len] for _ in range(batch_size)], dtype=torch.long, device=self.device
)
total_len = seq_len + 1
for tid in task_ids:
pool.task_extend(tid, seq_len)
input_mask = position_ids[:, :, None] >= torch.arange(
total_len, device=self.device
)
kv_cache = pool.bind_tasks(task_ids, [seq_len + 1] * batch_size, self.device)
kv_cache = pool.bind_tasks(task_ids, workspace, self.device)
with torch.inference_mode(), attn_backend(self.backend):
self.model(
input_ids,
@@ -123,6 +164,7 @@ class GenerationBenchmark:
import time
pool = self._make_pool(batch_size, prompt_length)
workspace = self._make_workspace(pool, self.config)
task_ids = [f"bench_prefill_{i}" for i in range(batch_size)]
for tid in task_ids:
pool.task_alloc(tid, list(range(prompt_length)))
@@ -138,9 +180,7 @@ class GenerationBenchmark:
input_mask = position_ids.unsqueeze(-1) >= torch.arange(
prompt_length, device=self.device
)
kv_cache = pool.bind_tasks(
task_ids, [prompt_length] * batch_size, self.device, start_pos=0
)
kv_cache = pool.bind_tasks(task_ids, workspace, self.device, start_pos=0)
for _ in range(3):
with torch.inference_mode(), attn_backend(self.backend):
@@ -180,22 +220,109 @@ class GenerationBenchmark:
prompt_length: int = 512,
gen_length: int = 128,
num_trials: int = 5,
) -> BenchmarkResult:
if self.cuda_graph and self.backend == "cuda":
return self._run_graph_decode_benchmark(
batch_size, prompt_length, gen_length, num_trials
)
return self._run_plain_decode_benchmark(
batch_size, prompt_length, gen_length, num_trials
)
def _run_graph_decode_benchmark(
self,
batch_size: int,
prompt_length: int,
gen_length: int,
num_trials: int,
) -> BenchmarkResult:
import time
# Decode grows seq_len monotonically up to prompt + 5 + gen*num_trials
# (warmup 5 steps, then one step per trial), so size the pool to cover it.
max_seq_len = prompt_length + 5 + gen_length * num_trials
pool = self._make_pool(batch_size, max_seq_len)
task_ids = self._run_prefill(pool, batch_size, prompt_length)
workspace = self._make_workspace(pool, self.config)
task_ids = self._run_prefill(pool, batch_size, prompt_length, workspace)
b = batch_size
input_ids_buf = torch.zeros(b, 1, dtype=torch.long, device=self.device)
position_ids_buf = torch.zeros(b, dtype=torch.long, device=self.device)
arange = torch.arange(max_seq_len, device=self.device)
gctx = CudaGraphContext(enabled=True)
graph_key = (b,)
def _decode_graph_step(seq_len):
input_ids_buf.copy_(
torch.randint(0, self.config.vocab_size, (b, 1), device=self.device)
)
position_ids_buf[:] = seq_len
for tid in task_ids:
pool.task_extend(tid, seq_len)
kv_cache = pool.bind_tasks(task_ids, workspace, self.device)
input_mask = torch.ge(
position_ids_buf[:, None],
arange,
out=workspace.input_mask[:b, 0, :max_seq_len],
)
input_mask = input_mask.unsqueeze(1)
with torch.inference_mode(), attn_backend(self.backend):
return gctx.forward(
self.model,
key=graph_key,
input_ids=input_ids_buf,
input_mask=input_mask,
kv_cache=kv_cache,
position_ids=position_ids_buf.unsqueeze(1),
)
for i in range(5):
self._run_decode_step(pool, task_ids, prompt_length + i)
_decode_graph_step(prompt_length + i)
torch.cuda.synchronize()
t0 = time.perf_counter()
for i in range(gen_length * num_trials):
self._run_decode_step(pool, task_ids, prompt_length + 5 + i)
_decode_graph_step(prompt_length + 5 + i)
torch.cuda.synchronize()
elapsed = time.perf_counter() - t0
tokens = batch_size * gen_length * num_trials
tps = tokens / elapsed
return BenchmarkResult(
name="decode",
batch_size=batch_size,
seq_len=gen_length,
tokens_per_second=tps,
latency_ms=elapsed / (gen_length * num_trials) * 1000,
metadata={
"benchmark_type": "decode",
"num_trials": num_trials,
"prompt_length": prompt_length,
"cuda_graph": True,
},
)
def _run_plain_decode_benchmark(
self,
batch_size: int,
prompt_length: int,
gen_length: int,
num_trials: int,
) -> BenchmarkResult:
import time
max_seq_len = prompt_length + 5 + gen_length * num_trials
pool = self._make_pool(batch_size, max_seq_len)
workspace = self._make_workspace(pool, self.config)
task_ids = self._run_prefill(pool, batch_size, prompt_length, workspace)
for i in range(5):
self._run_decode_step(pool, task_ids, prompt_length + i, workspace)
torch.cuda.synchronize()
t0 = time.perf_counter()
for i in range(gen_length * num_trials):
self._run_decode_step(pool, task_ids, prompt_length + 5 + i, workspace)
torch.cuda.synchronize()
elapsed = time.perf_counter() - t0
tokens = batch_size * gen_length * num_trials
@@ -250,11 +377,27 @@ def print_benchmark_result(result: BenchmarkResult) -> None:
@click.option("--num_trials", type=int, default=5, help="Number of trials.")
@click.option("--prefill_only", is_flag=True, help="Prefill benchmark only.")
@click.option("--decode_only", is_flag=True, help="Decode benchmark only.")
@click.option(
"--cuda-graph",
is_flag=True,
help="Enable CUDA graph capture for decode (cuda backend only).",
)
@click.option(
"--ckpt",
required=True,
required=False,
default=None,
type=click.Path(exists=True, file_okay=False, dir_okay=True, path_type=Path),
help="Checkpoint directory.",
help="Checkpoint directory. If omitted, a randomly-initialized model is "
"built from --config or the default 1B GQA preset.",
)
@click.option(
"--config",
"config_path",
required=False,
default=None,
type=click.Path(exists=True, file_okay=True, dir_okay=False, path_type=Path),
help="Optional model config JSON (used when --ckpt is omitted to define the "
"architecture). Defaults to the 1B GQA preset.",
)
def benchmark_command(
device: str,
@@ -268,7 +411,9 @@ def benchmark_command(
num_trials: int,
prefill_only: bool,
decode_only: bool,
ckpt: str,
cuda_graph: bool,
ckpt: Optional[str],
config_path: Optional[Path],
) -> None:
"""Benchmark model throughput and latency."""
dtype_map: dict[str, torch.dtype] = {
@@ -277,9 +422,23 @@ def benchmark_command(
"float32": torch.float32,
}
click.echo(f"Loading model from {ckpt} ...")
config = AutoRegressiveLMConfig.from_file(str(Path(ckpt) / "config.json"))
model = AutoModel.from_pretrained(ckpt)
if ckpt is not None:
click.echo(f"Loading model from {ckpt} ...")
config = ConfigFactory.load(
json.loads((Path(ckpt) / "config.json").read_text(encoding="utf-8-sig"))
)
model = AutoModel.from_pretrained(ckpt)
else:
raw = dict(_DEFAULT_CONFIG)
if config_path is not None:
raw.update(json.loads(config_path.read_text(encoding="utf-8-sig")))
config = ConfigFactory.load(raw)
model = AutoRegressiveLM(config)
click.echo(
f"Using randomly-initialized model "
f"({sum(p.numel() for p in model.parameters()) / 1e9:.2f}B params)"
)
model.to(device=device, dtype=dtype_map[dtype])
model.eval()
@@ -293,6 +452,7 @@ def benchmark_command(
dtype=dtype_map[dtype],
cache_type=cache,
backend=name,
cuda_graph=cuda_graph,
)
click.secho(
-1
View File
@@ -6,7 +6,6 @@ import numpy as np
import pytest
import torch
from astrai.config.preprocess_config import PipelineConfig
from astrai.dataset.dataset import (
DatasetFactory,
GRPODataset,
+1 -1
View File
@@ -5,7 +5,7 @@ import torch
from astrai.config.model_config import AutoRegressiveLMConfig
from astrai.model.transformer import AutoRegressiveLM
from tests.conftest import skip_no_kernel
from tests.conftest import skip_no_kernel # noqa: F401 re-export for test modules
D = 64
CFG = dict(
+20 -7
View File
@@ -10,27 +10,37 @@ from astrai.extension import (
ATTN_BACKEND,
AttentionBackendFactory,
CudaBackend,
TorchNativeBackend,
attn_backend,
get_backend,
)
def test_default_backend_is_torch_native():
"""Default is the highest-priority available backend (flash > cuda > torch)."""
from astrai.extension.attention_backend import (
CudaBackend,
FlashAttnBackend,
TorchNativeBackend,
_resolve_default_backend,
)
backend = get_backend()
assert isinstance(backend, TorchNativeBackend)
assert isinstance(backend, (CudaBackend, FlashAttnBackend, TorchNativeBackend))
assert isinstance(backend, type(_resolve_default_backend()))
def test_attn_backend_context_with_enum():
default = get_backend()
with attn_backend(ATTN_BACKEND.CUDA):
assert isinstance(get_backend(), CudaBackend)
assert isinstance(get_backend(), TorchNativeBackend)
assert get_backend() is default
def test_attn_backend_context_with_registered_name():
default = get_backend()
with attn_backend("cuda"):
assert isinstance(get_backend(), CudaBackend)
assert isinstance(get_backend(), TorchNativeBackend)
assert get_backend() is default
def test_attention_backend_factory_lists_builtin_backends():
@@ -48,19 +58,22 @@ def test_attn_backend_rejects_unknown_registered_name():
def test_attn_backend_context_with_class():
default = get_backend()
with attn_backend(CudaBackend):
assert isinstance(get_backend(), CudaBackend)
assert isinstance(get_backend(), TorchNativeBackend)
assert get_backend() is default
def test_attn_backend_context_with_instance():
custom = CudaBackend()
default = get_backend()
with attn_backend(custom):
assert get_backend() is custom
assert isinstance(get_backend(), TorchNativeBackend)
assert get_backend() is default
def test_cudabackend_is_context_manager():
default = get_backend()
with CudaBackend():
assert isinstance(get_backend(), CudaBackend)
assert isinstance(get_backend(), TorchNativeBackend)
assert get_backend() is default
+15 -9
View File
@@ -14,7 +14,12 @@ from tests.extension.conftest import D, skip_no_kernel
def _ws(pool: PagePool) -> InferenceWorkspace:
return InferenceWorkspace(
pool.max_batch_size, pool.max_seq_len, pool.device, pool.dtype
pool.max_batch_size,
pool.max_seq_len,
max_q_heads=2,
head_dim=64,
device=pool.device,
dtype=pool.dtype,
)
@@ -22,11 +27,10 @@ def _ws(pool: PagePool) -> InferenceWorkspace:
def test_training_forward_matches_torch(cuda_model):
"""Training forward (kv_cache=None) should produce identical logits.
CudaBackend is inference-only: it raises when kv_cache is None. Training
must use TorchNativeBackend (the default). Verify the torch path is
stable and that CudaBackend rejects the training path explicitly.
CudaBackend is now safe as a default: for training (``kv_cache=None``)
or non-bf16 inputs it falls back to torch SDPA. Verify the fallback
path matches the torch-native forward exactly.
"""
import pytest
model, _ = cuda_model
input_ids = torch.randint(0, 1000, (2, 16), device="cuda")
@@ -34,11 +38,13 @@ def test_training_forward_matches_torch(cuda_model):
with torch.no_grad():
out_torch = model(input_ids)
with pytest.raises(RuntimeError, match="does not support training"):
with attn_backend(ATTN_BACKEND.CUDA):
with torch.no_grad():
model(input_ids)
with attn_backend(ATTN_BACKEND.CUDA):
with torch.no_grad():
out_cuda = model(input_ids)
torch.testing.assert_close(
out_cuda["logits"], out_torch["logits"], atol=1e-6, rtol=1e-6
)
assert out_torch["logits"].shape[0] == 2
+6 -1
View File
@@ -16,7 +16,12 @@ from astrai.inference.core.workspace import InferenceWorkspace
def _ws(pool: PagePool) -> InferenceWorkspace:
"""Workspace sized to the pool (bind_tasks requires it)."""
return InferenceWorkspace(
pool.max_batch_size, pool.max_seq_len, pool.device, pool.dtype
pool.max_batch_size,
pool.max_seq_len,
max_q_heads=2,
head_dim=4,
device=pool.device,
dtype=pool.dtype,
)
+27 -7
View File
@@ -4,7 +4,7 @@ import threading
from unittest.mock import MagicMock, patch
from astrai.inference import STOP
from astrai.inference.engine import GenerateResult
from astrai.inference.engine import GenerateResult, InferenceEngine
def test_result_append_single():
@@ -101,8 +101,6 @@ def test_result_get_results():
def test_engine_generate_non_streaming_single():
from astrai.inference.engine import InferenceEngine
mock_model = MagicMock()
mock_tokenizer = MagicMock()
mock_tokenizer.encode.return_value = [1, 2, 3]
@@ -126,8 +124,6 @@ def test_engine_generate_non_streaming_single():
def test_engine_generate_streaming_yields_tokens():
from astrai.inference.engine import InferenceEngine
mock_model = MagicMock()
mock_tokenizer = MagicMock()
mock_tokenizer.encode.return_value = [1, 2, 3]
@@ -157,8 +153,6 @@ def test_engine_generate_streaming_yields_tokens():
def test_engine_generate_non_streaming_batch():
from astrai.inference.engine import InferenceEngine
mock_model = MagicMock()
mock_tokenizer = MagicMock()
mock_tokenizer.encode.return_value = [1, 2, 3]
@@ -179,3 +173,29 @@ def test_engine_generate_non_streaming_batch():
eng = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=2)
results = eng.generate(["hello", "world"])
assert results == ["r", "r"]
def test_engine_generate_zero_max_tokens_returns_empty():
mock_model = MagicMock()
mock_tokenizer = MagicMock()
mock_tokenizer.encode.return_value = [1, 2, 3]
mock_tokenizer.stop_ids = [0]
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
instance = MockSched.return_value
instance.remove_task.return_value = []
eng = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=2)
assert eng.generate(["hello", "world"], max_tokens=0) == ["", ""]
instance.add_task.assert_not_called()
def test_engine_generate_zero_max_tokens_stream_is_empty():
mock_model = MagicMock()
mock_tokenizer = MagicMock()
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
instance = MockSched.return_value
eng = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=1)
assert list(eng.generate("hello", stream=True, max_tokens=0)) == []
instance.add_task.assert_not_called()
-30
View File
@@ -8,7 +8,6 @@ import pytest
from astrai.inference.api.anthropic import AnthropicResponseBuilder
from astrai.inference.api.openai import OpenAIResponseBuilder
from astrai.inference.api.protocol import GenContext, StopChecker, StopInfo
from astrai.inference.engine import GenerationRequest
def _make_ctx(**kwargs):
@@ -255,32 +254,3 @@ class TestAnthropicResponseBuilder:
resp = builder.format_response(ctx, "full text", stop)
assert resp["content"][0]["text"] == "full text"
assert resp["stop_reason"] == "end_turn"
class TestGenerationRequestValidation:
def test_valid_params(self):
gr = GenerationRequest(
messages=[{"role": "user", "content": "hi"}],
top_k=50,
top_p=0.9,
temperature=0.7,
)
assert gr.top_k == 50
def test_invalid_top_p_raises(self):
with pytest.raises(ValueError, match="top_p"):
GenerationRequest(messages=[{"role": "user", "content": "hi"}], top_p=1.5)
def test_invalid_top_k_raises(self):
with pytest.raises(ValueError, match="top_k"):
GenerationRequest(messages=[{"role": "user", "content": "hi"}], top_k=-1)
def test_invalid_temperature_raises(self):
with pytest.raises(ValueError, match="temperature"):
GenerationRequest(
messages=[{"role": "user", "content": "hi"}], temperature=-0.1
)
def test_top_k_zero_valid(self):
gr = GenerationRequest(messages=[{"role": "user", "content": "hi"}], top_k=0)
assert gr.top_k == 0
+8
View File
@@ -261,6 +261,14 @@ def test_run_batch_respects_max_tokens(device):
scheduler.stop()
def test_run_batch_zero_max_tokens_returns_empty(device):
scheduler, _tok, _model = _make_real_scheduler(device)
try:
assert scheduler.run_batch([[10, 20, 30]], max_tokens=0) == [[]]
finally:
scheduler.stop()
def test_run_batch_stop_id_terminates(device):
"""A token matching stop_ids terminates generation for that prompt."""
scheduler, _tok, _model = _make_real_scheduler(device)
+1 -5
View File
@@ -5,14 +5,12 @@ multi-rank environment without requiring multiple GPUs.
"""
import torch
import torch.distributed as dist
import torch.nn as nn
from astrai.model.transformer import AutoRegressiveLM
from astrai.parallel import get_rank, spawn_parallel_fn
from astrai.parallel.executor import broadcast_state_dict, create_ref_model
from astrai.trainer.strategy import GRPOStrategy
from tests.helpers import FakeExecutor, make_rollout_config
from tests.helpers import make_rollout_config
def _broadcast_worker():
@@ -77,8 +75,6 @@ def _create_ref_model_worker():
)
assert ref is not None, f"rank {rank}: ref model is None"
# Every rank should have rank-0's weights, not its own
rank0_sd = model.state_dict() if rank == 0 else None
# Broadcast rank-0's original weights for comparison
if rank == 0:
expected_sd = {k: v.clone() for k, v in model.state_dict().items()}
+2 -2
View File
@@ -71,7 +71,7 @@ def _inner_run(batch_per_device, ckpt_interval, ckpt_dir, ready_file):
optimizer_fn=optimizer_fn,
scheduler_fn=scheduler_fn,
ckpt_dir=ckpt_dir,
n_epoch=1,
n_epoch=99999,
batch_per_device=batch_per_device,
ckpt_interval=ckpt_interval,
grad_accum_steps=1,
@@ -94,7 +94,7 @@ def _spawn_train_and_signal(ckpt_dir, sig, timeout=120):
)
p.start()
deadline = time.time() + 30
deadline = time.time() + 10
while time.time() < deadline:
if os.path.exists(ready_file):
with open(ready_file) as f: