Files
AstrAI/tests/extension/test_backend.py
T
ViperEkura 47b3ed4e44 feat: propagate attention backend across scheduler threads
- InferenceEngine/Scheduler accept an explicit backend
- capture request-level attn_backend context onto Task
- split prefill/decode batches by backend instance
- ASTR_BACKEND env overrides ContextVar as process-wide policy
- report resolved backend and CUDA-graph state in benchmark
2026-08-09 13:32:40 +08:00

93 lines
2.6 KiB
Python

"""Backend selection and context-manager switching tests.
These tests do not require CUDA — they only check that the active
backend is correctly set and restored.
"""
import pytest
from astrai.extension import (
ATTN_BACKEND,
AttentionBackendFactory,
CudaBackend,
attn_backend,
get_backend,
)
def test_default_backend_resolves_to_available():
"""Default backend is the first available in cuda > flash > torch order."""
from astrai.extension.attention_backend import (
CudaBackend,
FlashAttnBackend,
TorchNativeBackend,
_resolve_default_backend,
)
backend = get_backend()
assert isinstance(backend, (CudaBackend, FlashAttnBackend, TorchNativeBackend))
def test_attn_backend_context_with_enum():
default = get_backend()
with attn_backend(ATTN_BACKEND.CUDA):
assert isinstance(get_backend(), CudaBackend)
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 get_backend() is default
def test_backend_can_read_only_context_selection():
assert get_backend(use_default=False) is None
with attn_backend("cuda") as backend:
assert get_backend(use_default=False) is backend
assert get_backend(use_default=False) is None
def test_environment_backend_overrides_context(monkeypatch):
monkeypatch.setenv("ASTR_BACKEND", "torch_native")
with attn_backend("cuda"):
assert type(get_backend()).__name__ == "TorchNativeBackend"
assert type(get_backend(use_default=False)).__name__ == "TorchNativeBackend"
def test_attention_backend_factory_lists_builtin_backends():
assert AttentionBackendFactory.list_registered() == [
"cuda",
"flash",
"torch_native",
]
def test_attn_backend_rejects_unknown_registered_name():
with pytest.raises(ValueError, match="Unknown component: 'unknown'"):
with attn_backend("unknown"):
pass
def test_attn_backend_context_with_class():
default = get_backend()
with attn_backend(CudaBackend):
assert isinstance(get_backend(), CudaBackend)
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 get_backend() is default
def test_cudabackend_is_context_manager():
default = get_backend()
with CudaBackend():
assert isinstance(get_backend(), CudaBackend)
assert get_backend() is default