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:
@@ -19,12 +19,13 @@ 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, (CudaBackend, TorchNativeBackend))
|
||||
assert isinstance(backend, (CudaBackend, FlashAttnBackend, TorchNativeBackend))
|
||||
assert isinstance(backend, type(_resolve_default_backend()))
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user