refactor : use factory for attention backends

- register built-in backends through BaseFactory
- derive benchmark choices from registered backends
- cover string selection and invalid backend names
This commit is contained in:
2026-08-05 15:37:22 +08:00
parent 8c052c99ee
commit 8152760b5f
4 changed files with 43 additions and 21 deletions
+21
View File
@@ -8,6 +8,7 @@ import pytest
from astrai.extension import (
ATTN_BACKEND,
AttentionBackendFactory,
CudaBackend,
TorchNativeBackend,
attn_backend,
@@ -26,6 +27,26 @@ def test_attn_backend_context_with_enum():
assert isinstance(get_backend(), TorchNativeBackend)
def test_attn_backend_context_with_registered_name():
with attn_backend("cuda"):
assert isinstance(get_backend(), CudaBackend)
assert isinstance(get_backend(), 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():
with attn_backend(CudaBackend):
assert isinstance(get_backend(), CudaBackend)