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
This commit is contained in:
2026-08-09 13:32:40 +08:00
parent c1d05ae11d
commit 47b3ed4e44
9 changed files with 334 additions and 110 deletions
+14
View File
@@ -42,6 +42,20 @@ def test_attn_backend_context_with_registered_name():
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",
+38
View File
@@ -3,6 +3,7 @@
import threading
from unittest.mock import MagicMock, patch
from astrai.extension import TorchNativeBackend, attn_backend
from astrai.inference import STOP
from astrai.inference.engine import GenerateResult, InferenceEngine
@@ -199,3 +200,40 @@ def test_engine_generate_zero_max_tokens_stream_is_empty():
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()
def test_engine_passes_backend_to_scheduler():
mock_model = MagicMock()
mock_tokenizer = MagicMock()
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
InferenceEngine(
mock_model,
mock_tokenizer,
max_batch_size=1,
backend="torch_native",
)
assert MockSched.call_args.kwargs["backend"] == "torch_native"
def test_generate_captures_calling_backend_context():
mock_model = MagicMock()
mock_tokenizer = MagicMock()
captured = []
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
instance = MockSched.return_value
def fake_add(prompt, **kwargs):
captured.append(kwargs["backend"])
kwargs["stream_callback"](STOP)
return "task"
instance.add_task.side_effect = fake_add
engine = InferenceEngine(mock_model, mock_tokenizer)
with attn_backend("torch_native"):
assert engine.generate("hello") == ""
assert len(captured) == 1
assert isinstance(captured[0], TorchNativeBackend)
+65
View File
@@ -7,7 +7,9 @@ from unittest.mock import MagicMock, patch
import pytest
import torch
from astrai.extension import CudaBackend, TorchNativeBackend, get_backend
from astrai.inference import InferenceScheduler
from astrai.inference.metrics import MetricsCollector
from astrai.inference.runtime.executor import DecodeSteadyState, Executor
from astrai.inference.task import Task
from astrai.model.transformer import AutoRegressiveLM
@@ -76,6 +78,69 @@ def test_scheduler_concurrent_add_task(mock_model_and_tokenizer):
assert len(results["task_ids"]) == 50
def test_generation_loop_activates_backend_in_worker_thread():
scheduler = object.__new__(InferenceScheduler)
scheduler._backend = TorchNativeBackend()
scheduler._stop_event = threading.Event()
scheduler._task_cache = MagicMock()
observed = []
task_mgr = MagicMock()
task_mgr.tokenizer.stop_ids = [0]
task_mgr.remove_finished_tasks.return_value = []
task_mgr.get_active_tasks.return_value = []
task_mgr.max_batch_size = 1
task_mgr.pull_candidates.return_value = []
task_mgr.has_work.return_value = False
def observe_backend(*args, **kwargs):
observed.append(type(get_backend()))
scheduler._stop_event.set()
task_mgr.wait_for_tasks.side_effect = observe_backend
scheduler._task_mgr = task_mgr
thread = threading.Thread(target=scheduler._run_generation_loop)
thread.start()
thread.join(timeout=5)
assert not thread.is_alive()
assert observed == [TorchNativeBackend]
def test_step_splits_decode_batch_by_request_backend():
scheduler = object.__new__(InferenceScheduler)
scheduler._task_cache = MagicMock()
scheduler._task_cache.task_extend.return_value = True
scheduler._metrics = MetricsCollector()
scheduler._executor = MagicMock()
observed = []
def execute(tasks, **kwargs):
observed.append((type(get_backend()), [task.task_id for task in tasks]))
return [1] * len(tasks)
scheduler._executor.execute_decode.side_effect = execute
torch_task = Task("torch", [1], backend=TorchNativeBackend())
cuda_task = Task("cuda", [1], backend=CudaBackend())
for task in (torch_task, cuda_task):
task.input_tokens = 1
task.output_ids = [1]
task.mark_prefill_done()
scheduler._metrics.register(task.task_id)
produced, aborted = scheduler._step([torch_task, cuda_task])
assert aborted == []
assert produced == [torch_task, cuda_task]
assert observed == [
(TorchNativeBackend, ["torch"]),
(CudaBackend, ["cuda"]),
]
def test_scheduler_concurrent_add_remove_task(mock_model_and_tokenizer):
"""Test concurrent add and remove task operations."""
mock_model, mock_tokenizer = mock_model_and_tokenizer