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()
This commit is contained in:
@@ -21,7 +21,6 @@ from astrai.dataset import (
|
|||||||
)
|
)
|
||||||
from astrai.factory import BaseFactory
|
from astrai.factory import BaseFactory
|
||||||
from astrai.inference import (
|
from astrai.inference import (
|
||||||
GenerationRequest,
|
|
||||||
InferenceEngine,
|
InferenceEngine,
|
||||||
ProtocolHandler,
|
ProtocolHandler,
|
||||||
SamplingPipeline,
|
SamplingPipeline,
|
||||||
@@ -98,7 +97,6 @@ __all__ = [
|
|||||||
"EmbeddingEncoder",
|
"EmbeddingEncoder",
|
||||||
"EncoderConfig",
|
"EncoderConfig",
|
||||||
"ExecutorFactory",
|
"ExecutorFactory",
|
||||||
"GenerationRequest",
|
|
||||||
"InferenceEngine",
|
"InferenceEngine",
|
||||||
"LoRAConfig",
|
"LoRAConfig",
|
||||||
"Pipeline",
|
"Pipeline",
|
||||||
|
|||||||
@@ -22,7 +22,8 @@ Usage — mirroring ``torch.nn.attention.sdpa_kernel``:
|
|||||||
|
|
||||||
Thread-safe via ``contextvars`` — each scheduler thread gets its own
|
Thread-safe via ``contextvars`` — each scheduler thread gets its own
|
||||||
active backend. ``get_backend()`` returns the active one, falling back
|
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]``
|
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]``.
|
(blhd). The backend returns ``[batch, seq_len, n_heads * head_dim]``.
|
||||||
@@ -30,8 +31,9 @@ Layout convention: all q/k/v are ``[batch, seq_len, n_heads, head_dim]``
|
|||||||
|
|
||||||
import contextvars
|
import contextvars
|
||||||
import enum
|
import enum
|
||||||
|
import functools
|
||||||
import importlib
|
import importlib
|
||||||
import threading
|
import os
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from typing import TYPE_CHECKING, Optional, Union
|
from typing import TYPE_CHECKING, Optional, Union
|
||||||
@@ -54,60 +56,15 @@ _current_backend: contextvars.ContextVar["AttentionBackend"] = contextvars.Conte
|
|||||||
"attn_backend"
|
"attn_backend"
|
||||||
)
|
)
|
||||||
|
|
||||||
_lock = threading.Lock()
|
|
||||||
_flash_available: Optional[bool] = None
|
|
||||||
|
|
||||||
|
|
||||||
|
@functools.lru_cache(maxsize=1)
|
||||||
def flash_attn_available() -> bool:
|
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():
|
if not torch.cuda.is_available():
|
||||||
return False
|
return False
|
||||||
fa = _get_flash_attn()
|
fa = _get_flash_attn()
|
||||||
if fa is None:
|
if fa is None:
|
||||||
return False
|
return False
|
||||||
|
|
||||||
# version + compute-capability gate:
|
|
||||||
# FlashAttention-2 kernels need sm_70+; FlashAttention-3 (tcgen05,
|
|
||||||
# sm_90/sm_100) needs sm_90+.
|
|
||||||
try:
|
try:
|
||||||
major = int(fa.__version__.split(".")[0])
|
major = int(fa.__version__.split(".")[0])
|
||||||
cc = torch.cuda.get_device_capability()
|
cc = torch.cuda.get_device_capability()
|
||||||
@@ -117,8 +74,6 @@ def _flash_attn_check() -> bool:
|
|||||||
if (major >= 3 and cc_num < 90) or (major < 3 and 0 < cc_num < 70):
|
if (major >= 3 and cc_num < 90) or (major < 3 and 0 < cc_num < 70):
|
||||||
return False
|
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:
|
try:
|
||||||
if not hasattr(fa, "flash_attn_func"):
|
if not hasattr(fa, "flash_attn_func"):
|
||||||
return False
|
return False
|
||||||
@@ -129,6 +84,14 @@ def _flash_attn_check() -> bool:
|
|||||||
return False
|
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):
|
class ATTN_BACKEND(enum.Enum):
|
||||||
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
|
"""Backend selector enum, mirroring ``torch.nn.attention.SDPBackend``."""
|
||||||
|
|
||||||
@@ -141,12 +104,12 @@ _default_backend: Optional["AttentionBackend"] = None
|
|||||||
|
|
||||||
|
|
||||||
def _priority_backends() -> list["AttentionBackend"]:
|
def _priority_backends() -> list["AttentionBackend"]:
|
||||||
"""Available backends in priority order: flash -> cuda -> torch."""
|
"""Available backends in priority order: cuda -> flash -> torch."""
|
||||||
backends: list[AttentionBackend] = []
|
backends: list[AttentionBackend] = []
|
||||||
if flash_attn_available():
|
|
||||||
backends.append(FlashAttnBackend())
|
|
||||||
if is_available("attn_paged_decode") and is_available("attn_paged_prefill"):
|
if is_available("attn_paged_decode") and is_available("attn_paged_prefill"):
|
||||||
backends.append(CudaBackend())
|
backends.append(CudaBackend())
|
||||||
|
if flash_attn_available():
|
||||||
|
backends.append(FlashAttnBackend())
|
||||||
backends.append(TorchNativeBackend())
|
backends.append(TorchNativeBackend())
|
||||||
return backends
|
return backends
|
||||||
|
|
||||||
@@ -179,20 +142,31 @@ def _backend_supports(
|
|||||||
|
|
||||||
|
|
||||||
def _resolve_default_backend() -> "AttentionBackend":
|
def _resolve_default_backend() -> "AttentionBackend":
|
||||||
"""Pick the highest-priority available backend: flash -> cuda -> torch.
|
"""Pick the highest-priority available backend (cuda -> flash -> torch).
|
||||||
|
|
||||||
Resolved lazily on first ``get_backend()`` (flash/cuda availability is
|
Set ``ASTR_BACKEND`` to override: ``ASTR_BACKEND=cuda``, ``torch_native``,
|
||||||
checked once and cached). Per-call capability fallback happens in
|
or ``flash``. The value is the registered name (same as the
|
||||||
``attention()``, so this default is safe for training and fp32 models.
|
``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]
|
return _priority_backends()[0]
|
||||||
|
|
||||||
|
|
||||||
def get_backend() -> "AttentionBackend":
|
def get_backend() -> "AttentionBackend":
|
||||||
"""Return the active backend for the current thread/context.
|
"""Return the active backend for the current thread/context.
|
||||||
|
|
||||||
Falls back to the highest-priority available backend (flash -> cuda ->
|
Falls back to the highest-priority available backend (cuda -> flash ->
|
||||||
torch_native) when no backend has been activated via ``with``.
|
torch_native) when no backend has been activated via ``with``. Set
|
||||||
|
``ASTR_BACKEND`` to override the default.
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
return _current_backend.get()
|
return _current_backend.get()
|
||||||
@@ -252,6 +226,28 @@ 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(
|
def attention(
|
||||||
q: Tensor,
|
q: Tensor,
|
||||||
k: Tensor,
|
k: Tensor,
|
||||||
@@ -371,6 +367,17 @@ class AttentionBackend(ABC):
|
|||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
"""Multi-token prefill or training forward."""
|
"""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]):
|
class AttentionBackendFactory(BaseFactory[AttentionBackend]):
|
||||||
"""Factory for registered attention backends."""
|
"""Factory for registered attention backends."""
|
||||||
@@ -427,24 +434,7 @@ class TorchNativeBackend(AttentionBackend):
|
|||||||
is_causal: bool = False,
|
is_causal: bool = False,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
if kv_cache is not None:
|
if kv_cache is not None:
|
||||||
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
k, v = _write_and_gather_kv(kv_cache, k, v, layer_id, q, attn_mask)
|
||||||
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]
|
|
||||||
|
|
||||||
n_rep = q.size(2) // k.size(2)
|
n_rep = q.size(2) // k.size(2)
|
||||||
if n_rep > 1:
|
if n_rep > 1:
|
||||||
@@ -462,9 +452,6 @@ class TorchNativeBackend(AttentionBackend):
|
|||||||
return out
|
return out
|
||||||
|
|
||||||
|
|
||||||
_default_backend = None
|
|
||||||
|
|
||||||
|
|
||||||
@AttentionBackendFactory.register(ATTN_BACKEND.CUDA.value)
|
@AttentionBackendFactory.register(ATTN_BACKEND.CUDA.value)
|
||||||
class CudaBackend(AttentionBackend):
|
class CudaBackend(AttentionBackend):
|
||||||
"""CUDA kernel backend with direct KV cache access.
|
"""CUDA kernel backend with direct KV cache access.
|
||||||
@@ -487,11 +474,16 @@ class CudaBackend(AttentionBackend):
|
|||||||
def supports(**kwargs) -> bool:
|
def supports(**kwargs) -> bool:
|
||||||
head_dim = kwargs.get("head_dim", -1)
|
head_dim = kwargs.get("head_dim", -1)
|
||||||
return (
|
return (
|
||||||
head_dim in (32, 64, 128, 256)
|
torch.cuda.is_available()
|
||||||
|
and head_dim in (32, 64, 128, 256)
|
||||||
and is_available("attn_paged_decode")
|
and is_available("attn_paged_decode")
|
||||||
and is_available("attn_paged_prefill")
|
and is_available("attn_paged_prefill")
|
||||||
)
|
)
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def supports_graph() -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
def fwd_decode(
|
def fwd_decode(
|
||||||
self,
|
self,
|
||||||
q: Tensor,
|
q: Tensor,
|
||||||
@@ -572,12 +564,6 @@ class CudaBackend(AttentionBackend):
|
|||||||
return out.reshape(b, q_len, q.size(2), q.size(3)).flatten(2)
|
return out.reshape(b, q_len, q.size(2), q.size(3)).flatten(2)
|
||||||
|
|
||||||
|
|
||||||
def _kv_cache_is_contiguous(kv_cache: "KVCache") -> bool:
|
|
||||||
return kv_cache.k_buffer.size(1) == kv_cache.req_to_token.size(
|
|
||||||
0
|
|
||||||
) * kv_cache.req_to_token.size(1)
|
|
||||||
|
|
||||||
|
|
||||||
@AttentionBackendFactory.register(ATTN_BACKEND.FLASH.value)
|
@AttentionBackendFactory.register(ATTN_BACKEND.FLASH.value)
|
||||||
class FlashAttnBackend(AttentionBackend):
|
class FlashAttnBackend(AttentionBackend):
|
||||||
"""FlashAttention backend via the optional ``flash-attn`` package.
|
"""FlashAttention backend via the optional ``flash-attn`` package.
|
||||||
@@ -629,24 +615,11 @@ class FlashAttnBackend(AttentionBackend):
|
|||||||
is_causal: bool = False,
|
is_causal: bool = False,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
if kv_cache is not None:
|
if kv_cache is not None:
|
||||||
if q.size(1) == 1 and _kv_cache_is_contiguous(kv_cache):
|
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)
|
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)
|
||||||
kv_cache.k_buffer[layer_id, kv_cache.out_cache_loc] = k
|
|
||||||
kv_cache.v_buffer[layer_id, kv_cache.out_cache_loc] = v
|
|
||||||
|
|
||||||
max_len = kv_cache.max_len
|
|
||||||
indices = kv_cache.req_to_token[kv_cache.req_pool_indices, :max_len]
|
|
||||||
if q.size(1) == 1 and attn_mask is not None and attn_mask.dim() == 4:
|
|
||||||
pos_mask = attn_mask[:, 0, 0]
|
|
||||||
else:
|
|
||||||
pos_mask = (
|
|
||||||
torch.arange(max_len, device=q.device)[None, :]
|
|
||||||
< kv_cache.seq_lens[:, None]
|
|
||||||
)
|
|
||||||
indices = torch.where(pos_mask, indices, torch.zeros_like(indices))
|
|
||||||
k = kv_cache.k_buffer[layer_id, indices]
|
|
||||||
v = kv_cache.v_buffer[layer_id, indices]
|
|
||||||
|
|
||||||
n_rep = q.size(2) // k.size(2)
|
n_rep = q.size(2) // k.size(2)
|
||||||
if n_rep > 1:
|
if n_rep > 1:
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ Layers:
|
|||||||
- api/: HTTP orchestration (ProtocolHandler, server)
|
- api/: HTTP orchestration (ProtocolHandler, server)
|
||||||
- protocols/: Response builders (OpenAI, Anthropic)
|
- protocols/: Response builders (OpenAI, Anthropic)
|
||||||
- transport/: SSE transport utilities
|
- transport/: SSE transport utilities
|
||||||
- engine.py: Facade (InferenceEngine), Value Object (GenerationRequest)
|
- engine.py: Facade (InferenceEngine)
|
||||||
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy, FrequencyPenaltyStrategy)
|
- sample.py: Strategy pattern (TemperatureStrategy, TopKStrategy, TopPStrategy, FrequencyPenaltyStrategy)
|
||||||
"""
|
"""
|
||||||
|
|
||||||
@@ -42,7 +42,7 @@ from astrai.inference.core import (
|
|||||||
TaskStatus,
|
TaskStatus,
|
||||||
page_hash,
|
page_hash,
|
||||||
)
|
)
|
||||||
from astrai.inference.engine import GenerationRequest, InferenceEngine
|
from astrai.inference.engine import InferenceEngine
|
||||||
from astrai.inference.sample import (
|
from astrai.inference.sample import (
|
||||||
BaseSamplingStrategy,
|
BaseSamplingStrategy,
|
||||||
FrequencyPenaltyStrategy,
|
FrequencyPenaltyStrategy,
|
||||||
@@ -55,7 +55,6 @@ from astrai.inference.sample import (
|
|||||||
|
|
||||||
__all__ = [
|
__all__ = [
|
||||||
"InferenceEngine",
|
"InferenceEngine",
|
||||||
"GenerationRequest",
|
|
||||||
"InferenceScheduler",
|
"InferenceScheduler",
|
||||||
"Executor",
|
"Executor",
|
||||||
"STOP",
|
"STOP",
|
||||||
|
|||||||
@@ -179,9 +179,7 @@ class Executor:
|
|||||||
max_q_heads = config.num_attention_heads
|
max_q_heads = config.num_attention_heads
|
||||||
head_dim = config.hidden_size // config.num_attention_heads
|
head_dim = config.hidden_size // config.num_attention_heads
|
||||||
self._head_dim = head_dim
|
self._head_dim = head_dim
|
||||||
self._graph_supported = CudaBackend.supports(
|
self._graph_supported = CudaBackend.supports(head_dim=head_dim)
|
||||||
head_dim=head_dim
|
|
||||||
) and "cuda" in str(self.device)
|
|
||||||
self._workspace = InferenceWorkspace(
|
self._workspace = InferenceWorkspace(
|
||||||
max_batch_size=kv_cache.max_batch_size,
|
max_batch_size=kv_cache.max_batch_size,
|
||||||
max_seq_len=kv_cache.max_seq_len,
|
max_seq_len=kv_cache.max_seq_len,
|
||||||
@@ -367,7 +365,7 @@ class Executor:
|
|||||||
use_graph = (
|
use_graph = (
|
||||||
self._graph_ctx.enabled
|
self._graph_ctx.enabled
|
||||||
and self._graph_supported
|
and self._graph_supported
|
||||||
and isinstance(get_backend(), CudaBackend)
|
and get_backend().supports_graph()
|
||||||
)
|
)
|
||||||
key = (b,)
|
key = (b,)
|
||||||
if use_graph:
|
if use_graph:
|
||||||
|
|||||||
+36
-155
@@ -64,44 +64,6 @@ class GenerateResult:
|
|||||||
return self.results.copy()
|
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:
|
class InferenceEngine:
|
||||||
"""Unified inference engine backed by continuous-batching scheduler."""
|
"""Unified inference engine backed by continuous-batching scheduler."""
|
||||||
|
|
||||||
@@ -152,28 +114,17 @@ class InferenceEngine:
|
|||||||
results = [""] * len(prompts)
|
results = [""] * len(prompts)
|
||||||
return results if is_batch else results[0]
|
return results if is_batch else results[0]
|
||||||
|
|
||||||
if stream:
|
return self._generate(
|
||||||
return self._generate_streaming(
|
prompts,
|
||||||
prompts,
|
is_batch,
|
||||||
is_batch,
|
stream,
|
||||||
max_tokens,
|
max_tokens,
|
||||||
temperature,
|
temperature,
|
||||||
top_p,
|
top_p,
|
||||||
top_k,
|
top_k,
|
||||||
frequency_penalty,
|
frequency_penalty,
|
||||||
rep_window,
|
rep_window,
|
||||||
)
|
)
|
||||||
else:
|
|
||||||
return self._generate_non_streaming(
|
|
||||||
prompts,
|
|
||||||
is_batch,
|
|
||||||
max_tokens,
|
|
||||||
temperature,
|
|
||||||
top_p,
|
|
||||||
top_k,
|
|
||||||
frequency_penalty,
|
|
||||||
rep_window,
|
|
||||||
)
|
|
||||||
|
|
||||||
def generate_async(
|
def generate_async(
|
||||||
self,
|
self,
|
||||||
@@ -185,9 +136,10 @@ class InferenceEngine:
|
|||||||
frequency_penalty: float = 0.0,
|
frequency_penalty: float = 0.0,
|
||||||
rep_window: int = 64,
|
rep_window: int = 64,
|
||||||
) -> AsyncGenerator[str, None]:
|
) -> AsyncGenerator[str, None]:
|
||||||
sync_gen = self._generate_streaming(
|
sync_gen = self._generate(
|
||||||
[prompt],
|
[prompt],
|
||||||
False,
|
False,
|
||||||
|
True,
|
||||||
max_tokens,
|
max_tokens,
|
||||||
temperature,
|
temperature,
|
||||||
top_p,
|
top_p,
|
||||||
@@ -199,51 +151,30 @@ class InferenceEngine:
|
|||||||
async def _agen():
|
async def _agen():
|
||||||
loop = asyncio.get_event_loop()
|
loop = asyncio.get_event_loop()
|
||||||
while True:
|
while True:
|
||||||
token = await loop.run_in_executor(None, self._next_token, sync_gen)
|
try:
|
||||||
if token is None:
|
token = await loop.run_in_executor(None, next, sync_gen)
|
||||||
|
except StopIteration:
|
||||||
break
|
break
|
||||||
yield token
|
yield token
|
||||||
|
|
||||||
return _agen()
|
return _agen()
|
||||||
|
|
||||||
@staticmethod
|
def _generate(
|
||||||
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(
|
|
||||||
self,
|
self,
|
||||||
prompts: List[str],
|
prompts: List[str],
|
||||||
|
is_batch: bool,
|
||||||
|
stream: bool,
|
||||||
max_tokens: Optional[int],
|
max_tokens: Optional[int],
|
||||||
temperature: float,
|
temperature: float,
|
||||||
top_p: float,
|
top_p: float,
|
||||||
top_k: int,
|
top_k: int,
|
||||||
frequency_penalty: float,
|
frequency_penalty: float,
|
||||||
rep_window: int,
|
rep_window: int,
|
||||||
) -> Tuple[GenerateResult, List[str]]:
|
) -> Union[Generator, str, List[str]]:
|
||||||
n = len(prompts)
|
n = len(prompts)
|
||||||
result = GenerateResult(count=n)
|
result = GenerateResult(count=n)
|
||||||
task_ids = []
|
task_ids = [
|
||||||
for i, p in enumerate(prompts):
|
self.scheduler.add_task(
|
||||||
cb = self._make_callback(result, i)
|
|
||||||
task_id = self.scheduler.add_task(
|
|
||||||
prompt=p,
|
prompt=p,
|
||||||
max_tokens=max_tokens,
|
max_tokens=max_tokens,
|
||||||
temperature=temperature,
|
temperature=temperature,
|
||||||
@@ -251,39 +182,23 @@ class InferenceEngine:
|
|||||||
top_k=top_k,
|
top_k=top_k,
|
||||||
frequency_penalty=frequency_penalty,
|
frequency_penalty=frequency_penalty,
|
||||||
rep_window=rep_window,
|
rep_window=rep_window,
|
||||||
stream_callback=cb,
|
stream_callback=lambda token, idx=i: result.append(token, idx),
|
||||||
)
|
)
|
||||||
task_ids.append(task_id)
|
for i, p in enumerate(prompts)
|
||||||
return result, task_ids
|
]
|
||||||
|
|
||||||
@staticmethod
|
if not stream:
|
||||||
def _make_callback(result: GenerateResult, idx: int):
|
try:
|
||||||
def cb(token):
|
result.wait_completion()
|
||||||
result.append(token, idx)
|
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
|
remaining = n
|
||||||
finished = [False] * n
|
finished = [False] * n
|
||||||
|
|
||||||
@@ -307,40 +222,6 @@ class InferenceEngine:
|
|||||||
|
|
||||||
return gen()
|
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]:
|
def get_stats(self) -> Dict[str, Any]:
|
||||||
return self.scheduler.get_stats()
|
return self.scheduler.get_stats()
|
||||||
|
|
||||||
|
|||||||
@@ -19,12 +19,13 @@ def test_default_backend_is_torch_native():
|
|||||||
"""Default is the highest-priority available backend (flash > cuda > torch)."""
|
"""Default is the highest-priority available backend (flash > cuda > torch)."""
|
||||||
from astrai.extension.attention_backend import (
|
from astrai.extension.attention_backend import (
|
||||||
CudaBackend,
|
CudaBackend,
|
||||||
|
FlashAttnBackend,
|
||||||
TorchNativeBackend,
|
TorchNativeBackend,
|
||||||
_resolve_default_backend,
|
_resolve_default_backend,
|
||||||
)
|
)
|
||||||
|
|
||||||
backend = get_backend()
|
backend = get_backend()
|
||||||
assert isinstance(backend, (CudaBackend, TorchNativeBackend))
|
assert isinstance(backend, (CudaBackend, FlashAttnBackend, TorchNativeBackend))
|
||||||
assert isinstance(backend, type(_resolve_default_backend()))
|
assert isinstance(backend, type(_resolve_default_backend()))
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -8,7 +8,6 @@ import pytest
|
|||||||
from astrai.inference.api.anthropic import AnthropicResponseBuilder
|
from astrai.inference.api.anthropic import AnthropicResponseBuilder
|
||||||
from astrai.inference.api.openai import OpenAIResponseBuilder
|
from astrai.inference.api.openai import OpenAIResponseBuilder
|
||||||
from astrai.inference.api.protocol import GenContext, StopChecker, StopInfo
|
from astrai.inference.api.protocol import GenContext, StopChecker, StopInfo
|
||||||
from astrai.inference.engine import GenerationRequest
|
|
||||||
|
|
||||||
|
|
||||||
def _make_ctx(**kwargs):
|
def _make_ctx(**kwargs):
|
||||||
@@ -255,32 +254,3 @@ class TestAnthropicResponseBuilder:
|
|||||||
resp = builder.format_response(ctx, "full text", stop)
|
resp = builder.format_response(ctx, "full text", stop)
|
||||||
assert resp["content"][0]["text"] == "full text"
|
assert resp["content"][0]["text"] == "full text"
|
||||||
assert resp["stop_reason"] == "end_turn"
|
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
|
|
||||||
|
|||||||
Reference in New Issue
Block a user