perf: batch decode stream callbacks into one dispatch per step

- add BatchedStreamCallback sink type: TaskManager resolves a decode step's (task_id, token) events under one lock and delivers each sink a single list instead of one call per token
- keep the plain Callable[[str]] callback contract: per-token callbacks still receive one call per event, and invoke_callback/cancel_task wrap single events for batched sinks
- collect aborted, text, and finish STOP events in the scheduler decode loop and dispatch once per step instead of once per token
- register one _ResultSink per generate call (replacing per-task closures) so GenerateResult takes its lock and wakes waiters once per step, with late-bind replay for tasks that start decoding before add_task returns their id
- apply GenerateResult batches under a single condition hold via append_batch; append delegates to it
- update engine test fakes to the batched contract and add coverage for event grouping, single-event dispatch, cancel STOP, and late-bind replay

Benchmark: NVIDIA L20 (idle), CUDA 12.8, torch 2.11.0+cu128, 1.2B bf16 checkpoint, prompt 512, 256 greedy tokens, CUDA graph on, serving-level decode, 3 trials
- batch 32: 7.808 -> 7.506 ms/token (4098 -> 4263 batch tok/s, +4.0%)
- batch 1/8: unchanged within noise (3.768 -> 3.797 / 4.699 -> 4.607 ms/token)
- full suite: 896 passed
This commit is contained in:
2026-09-05 00:03:36 +08:00
parent 1798474316
commit 074642b6d2
6 changed files with 265 additions and 46 deletions
+9 -1
View File
@@ -17,12 +17,20 @@ from astrai.inference.network import get_app, run_server
from astrai.inference.runtime.executor import Executor from astrai.inference.runtime.executor import Executor
from astrai.inference.runtime.sample import sample from astrai.inference.runtime.sample import sample
from astrai.inference.scheduler import InferenceScheduler from astrai.inference.scheduler import InferenceScheduler
from astrai.inference.task import STOP, GenerationResult, Task, TaskManager, TaskStatus from astrai.inference.task import (
STOP,
BatchedStreamCallback,
GenerationResult,
Task,
TaskManager,
TaskStatus,
)
__all__ = [ __all__ = [
"InferenceEngine", "InferenceEngine",
"build_engine", "build_engine",
"InferenceScheduler", "InferenceScheduler",
"BatchedStreamCallback",
"GenerationResult", "GenerationResult",
"Executor", "Executor",
"STOP", "STOP",
+67 -10
View File
@@ -13,7 +13,7 @@ import torch.nn as nn
from astrai.extension import ATTN_BACKEND, AttentionBackend, get_backend from astrai.extension import ATTN_BACKEND, AttentionBackend, get_backend
from astrai.inference.cache import PagePool from astrai.inference.cache import PagePool
from astrai.inference.scheduler import InferenceScheduler from astrai.inference.scheduler import InferenceScheduler
from astrai.inference.task import STOP from astrai.inference.task import STOP, BatchedStreamCallback
from astrai.model import AutoModel from astrai.model import AutoModel
from astrai.tokenize import AutoTokenizer from astrai.tokenize import AutoTokenizer
@@ -33,15 +33,27 @@ class GenerateResult:
self._total = count self._total = count
def append(self, token: str, idx: int = 0): def append(self, token: str, idx: int = 0):
self.append_batch([(idx, token)])
def append_batch(self, items: List[Tuple[int, Any]]) -> None:
"""Append multiple ``(idx, token)`` events under one lock/notify.
Batched counterpart to :meth:`append` for per-step delivery: state
updates for every event happen under a single condition hold and
waiters are woken once per batch instead of once per token.
"""
if not items:
return
with self._cond: with self._cond:
for idx, token in items:
self.tokens.append((idx, token)) self.tokens.append((idx, token))
if token is not STOP: if token is STOP:
self.results[idx] += token
else:
if not self._done[idx]: if not self._done[idx]:
self._done[idx] = True self._done[idx] = True
self._completed += 1 self._completed += 1
self._cond.notify_all() self._cond.notify_all()
else:
self.results[idx] += token
self._event.set() self._event.set()
def pop_all(self) -> List[Tuple[int, str]]: def pop_all(self) -> List[Tuple[int, str]]:
@@ -69,6 +81,47 @@ class GenerateResult:
return self.results.copy() return self.results.copy()
class _ResultSink(BatchedStreamCallback):
"""Batched stream channel from the scheduler into one GenerateResult.
Registered as the ``stream_callback`` for every task of a single
``generate`` call, so the scheduler's one dispatch per decode step
maps to one ``append_batch`` (one lock, one waiter wake). A task can
start decoding the moment ``add_task`` returns — before the engine
learns its id — so events for ids not yet bound are buffered and
replayed on ``bind``.
"""
def __init__(self, result: GenerateResult):
self._result = result
self._lock = threading.Lock()
self._index_of: Dict[str, int] = {}
self._pending: List[Tuple[str, Any]] = []
def bind(self, task_id: str, idx: int) -> None:
with self._lock:
self._index_of[task_id] = idx
replay = [(idx, token) for tid, token in self._pending if tid == task_id]
if replay:
self._pending = [
(tid, token) for tid, token in self._pending if tid != task_id
]
if replay:
self._result.append_batch(replay)
def __call__(self, events: List[Tuple[str, Any]]) -> None:
with self._lock:
items: List[Tuple[int, Any]] = []
for tid, token in events:
idx = self._index_of.get(tid)
if idx is None:
self._pending.append((tid, token))
else:
items.append((idx, token))
if items:
self._result.append_batch(items)
class InferenceEngine: class InferenceEngine:
"""Unified inference engine backed by continuous-batching scheduler.""" """Unified inference engine backed by continuous-batching scheduler."""
@@ -147,6 +200,7 @@ class InferenceEngine:
) -> AsyncGenerator[str, None]: ) -> AsyncGenerator[str, None]:
request_backend = get_backend(use_default=False) request_backend = get_backend(use_default=False)
result = GenerateResult() result = GenerateResult()
sink = _ResultSink(result)
task_id = self.scheduler.add_task( task_id = self.scheduler.add_task(
prompt=prompt, prompt=prompt,
max_tokens=max_tokens, max_tokens=max_tokens,
@@ -156,8 +210,9 @@ class InferenceEngine:
frequency_penalty=frequency_penalty, frequency_penalty=frequency_penalty,
rep_window=rep_window, rep_window=rep_window,
backend=request_backend, backend=request_backend,
stream_callback=result.append, stream_callback=sink,
) )
sink.bind(task_id, 0)
async def _agen(): async def _agen():
finished = False finished = False
@@ -191,8 +246,10 @@ class InferenceEngine:
n = len(prompts) n = len(prompts)
request_backend = get_backend(use_default=False) request_backend = get_backend(use_default=False)
result = GenerateResult(count=n) result = GenerateResult(count=n)
task_ids = [ sink = _ResultSink(result)
self.scheduler.add_task( task_ids = []
for i, p in enumerate(prompts):
task_id = self.scheduler.add_task(
prompt=p, prompt=p,
max_tokens=max_tokens, max_tokens=max_tokens,
temperature=temperature, temperature=temperature,
@@ -201,10 +258,10 @@ class InferenceEngine:
frequency_penalty=frequency_penalty, frequency_penalty=frequency_penalty,
rep_window=rep_window, rep_window=rep_window,
backend=request_backend, backend=request_backend,
stream_callback=lambda token, idx=i: result.append(token, idx), stream_callback=sink,
) )
for i, p in enumerate(prompts) sink.bind(task_id, i)
] task_ids.append(task_id)
if not stream: if not stream:
try: try:
+7 -5
View File
@@ -264,17 +264,19 @@ class InferenceScheduler:
decoded, aborted = self._stepper.step(active) decoded, aborted = self._stepper.step(active)
for t in aborted: # One dispatch per step: batch-aware sinks take their
self._task_mgr.invoke_callback(t.task_id, STOP) # lock (and wake waiters) once instead of once per token.
events: List[Tuple[str, Any]] = [(t.task_id, STOP) for t in aborted]
for t in decoded: for t in decoded:
if t.status == TaskStatus.ABORTED: if t.status == TaskStatus.ABORTED:
continue continue
new_text = t.decode_new_token(self._task_mgr.tokenizer) new_text = t.decode_new_token(self._task_mgr.tokenizer)
if new_text: if new_text:
self._task_mgr.invoke_callback(t.task_id, new_text) events.append((t.task_id, new_text))
if t.is_finished(stop_ids): if t.is_finished(stop_ids):
self._task_mgr.invoke_callback(t.task_id, STOP) events.append((t.task_id, STOP))
if events:
self._task_mgr.invoke_callbacks(events)
except Exception as e: except Exception as e:
self._stop_event.set() self._stop_event.set()
+52 -2
View File
@@ -1,6 +1,7 @@
import threading import threading
import time import time
import uuid import uuid
from abc import ABC, abstractmethod
from collections import deque from collections import deque
from dataclasses import dataclass from dataclasses import dataclass
from enum import Enum from enum import Enum
@@ -145,6 +146,22 @@ class Task:
return False return False
class BatchedStreamCallback(ABC):
"""Stream sink that receives a whole scheduler step's events in one call.
The scheduling loop dispatches once per decode step: every
``(task_id, token)`` event routed to the same sink object is delivered
as a single list, so batch-aware consumers take their lock and wake
waiters once per step instead of once per token. Plain per-token
callbacks keep the ``Callable[[str], None]`` contract.
"""
@abstractmethod
def __call__(self, events: List[Tuple[str, Any]]) -> None:
"""Consume ``[(task_id, token), ...]`` produced by one decode step."""
raise NotImplementedError
class TaskManager: class TaskManager:
"""Thread-safe task queues and lifecycle transitions (no page ops).""" """Thread-safe task queues and lifecycle transitions (no page ops)."""
@@ -256,6 +273,9 @@ class TaskManager:
immediate = [task] immediate = [task]
if cancelled and callback is not None: if cancelled and callback is not None:
if isinstance(callback, BatchedStreamCallback):
callback([(task_id, STOP)])
else:
callback(STOP) callback(STOP)
return immediate, cancelled return immediate, cancelled
@@ -264,10 +284,40 @@ class TaskManager:
immediate, _ = self.cancel_task(task_id) immediate, _ = self.cancel_task(task_id)
return immediate return immediate
def invoke_callback(self, task_id: str, token: str): def invoke_callback(self, task_id: str, token: Any):
with self._lock: with self._lock:
cb = self._callbacks.get(task_id) cb = self._callbacks.get(task_id)
if cb: if isinstance(cb, BatchedStreamCallback):
cb([(task_id, token)])
elif cb:
cb(token)
def invoke_callbacks(self, events: List[Tuple[str, Any]]) -> None:
"""Dispatch one decode step's ``(task_id, token)`` events.
Callbacks resolve under a single lock acquisition; events aimed at
the same batched sink are delivered as one list (one consumer-side
lock/notify per step), while plain per-token callbacks receive one
call per event.
"""
grouped: Dict[int, Tuple[BatchedStreamCallback, List[Any]]] = {}
plain: List[Tuple[Callable[[str], None], Any]] = []
with self._lock:
for task_id, token in events:
cb = self._callbacks.get(task_id)
if cb is None:
continue
if isinstance(cb, BatchedStreamCallback):
entry = grouped.get(id(cb))
if entry is None:
grouped[id(cb)] = (cb, [(task_id, token)])
else:
entry[1].append((task_id, token))
else:
plain.append((cb, token))
for cb, batch in grouped.values():
cb(batch)
for cb, token in plain:
cb(token) cb(token)
def get_stats(self) -> Dict[str, Any]: def get_stats(self) -> Dict[str, Any]:
+62 -19
View File
@@ -1,6 +1,7 @@
"""Unit tests for GenerateResult accumulator and InferenceEngine.generate().""" """Unit tests for GenerateResult accumulator and InferenceEngine.generate()."""
import asyncio import asyncio
import itertools
import threading import threading
from unittest.mock import MagicMock, patch from unittest.mock import MagicMock, patch
@@ -8,7 +9,12 @@ import pytest
from astrai.extension import TorchNativeBackend, attn_backend from astrai.extension import TorchNativeBackend, attn_backend
from astrai.inference import STOP from astrai.inference import STOP
from astrai.inference.engine import GenerateResult, InferenceEngine, build_engine from astrai.inference.engine import (
GenerateResult,
InferenceEngine,
_ResultSink,
build_engine,
)
from tests.helpers import FakeTokenizer, make_model from tests.helpers import FakeTokenizer, make_model
@@ -50,6 +56,33 @@ def test_result_stop_does_not_double_count():
assert r._completed == 1 assert r._completed == 1
def test_result_append_batch_updates_state_in_one_commit():
r = GenerateResult(count=2)
r.append_batch([(0, "he"), (1, "wo"), (0, "llo"), (1, "rld")])
r.append_batch([(0, STOP), (1, STOP)])
assert r.results == ["hello", "world"]
assert r._completed == 2
assert r.pop_all() == [
(0, "he"),
(1, "wo"),
(0, "llo"),
(1, "rld"),
(0, STOP),
(1, STOP),
]
def test_result_sink_replays_events_arriving_before_bind():
r = GenerateResult(count=1)
sink = _ResultSink(r)
sink([("t0", "he")]) # task id not bound yet: buffered, not applied
assert r.results == [""]
sink.bind("t0", 0)
sink([("t0", "llo"), ("t0", STOP)])
assert r.results == ["hello"]
assert r._completed == 1
def test_result_pop_all_returns_and_clears(): def test_result_pop_all_returns_and_clears():
r = GenerateResult(count=2) r = GenerateResult(count=2)
r.append("a", 0) r.append("a", 0)
@@ -118,8 +151,8 @@ def test_engine_generate_non_streaming_single():
def fake_add(prompt, **kw): def fake_add(prompt, **kw):
cb = kw["stream_callback"] cb = kw["stream_callback"]
cb("response") cb([("task-1", "response"), ("task-1", STOP)])
cb(STOP) return "task-1"
instance.add_task.side_effect = fake_add instance.add_task.side_effect = fake_add
instance.remove_task.return_value = [] instance.remove_task.return_value = []
@@ -136,6 +169,7 @@ def test_engine_generate_streaming_yields_tokens():
def capture_cb(prompt, **kw): def capture_cb(prompt, **kw):
callbacks_saved.append(kw.get("stream_callback")) callbacks_saved.append(kw.get("stream_callback"))
return "task-0"
with patch("astrai.inference.engine.InferenceScheduler") as MockSched: with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
instance = MockSched.return_value instance = MockSched.return_value
@@ -146,9 +180,9 @@ def test_engine_generate_streaming_yields_tokens():
gen = eng.generate("hello", stream=True) gen = eng.generate("hello", stream=True)
cb = callbacks_saved[0] cb = callbacks_saved[0]
cb("t1") cb([("task-0", "t1")])
cb("t2") cb([("task-0", "t2")])
cb(STOP) cb([("task-0", STOP)])
tokens = list(gen) tokens = list(gen)
assert tokens == ["t1", "t2"] assert tokens == ["t1", "t2"]
@@ -169,7 +203,7 @@ def test_engine_stream_close_cancels_unfinished_task():
engine = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=1) engine = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=1)
stream = engine.generate("hello", stream=True) stream = engine.generate("hello", stream=True)
callbacks_saved[0]("t1") callbacks_saved[0]([("task-1", "t1")])
assert next(stream) == "t1" assert next(stream) == "t1"
stream.close() stream.close()
@@ -182,6 +216,7 @@ def test_engine_generate_async_yields_tokens_until_stop():
def capture_cb(prompt, **kw): def capture_cb(prompt, **kw):
callbacks_saved.append(kw.get("stream_callback")) callbacks_saved.append(kw.get("stream_callback"))
return "task-0"
with patch("astrai.inference.engine.InferenceScheduler") as MockSched: with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
instance = MockSched.return_value instance = MockSched.return_value
@@ -198,9 +233,9 @@ def test_engine_generate_async_yields_tokens_until_stop():
return out return out
cb = callbacks_saved[0] cb = callbacks_saved[0]
cb("t1") cb([("task-0", "t1")])
cb("t2") cb([("task-0", "t2")])
cb(STOP) cb([("task-0", STOP)])
assert asyncio.run(collect()) == ["t1", "t2"] assert asyncio.run(collect()) == ["t1", "t2"]
@@ -219,7 +254,7 @@ def test_engine_async_close_cancels_unfinished_task():
engine = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=1) engine = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=1)
stream = engine.generate_async("hello") stream = engine.generate_async("hello")
callbacks_saved[0]("t1") callbacks_saved[0]([("task-1", "t1")])
async def consume_then_close(): async def consume_then_close():
assert await anext(stream) == "t1" assert await anext(stream) == "t1"
@@ -233,20 +268,25 @@ def test_engine_async_close_cancels_unfinished_task():
def test_engine_generate_non_streaming_batch(): def test_engine_generate_non_streaming_batch():
mock_model, mock_tokenizer = _make_engine_mocks(decode="r") mock_model, mock_tokenizer = _make_engine_mocks(decode="r")
with patch("astrai.inference.engine.InferenceScheduler") as MockSched: counter = itertools.count()
instance = MockSched.return_value task_ids = []
def fake_add(prompt, **kw): def fake_add(prompt, **kw):
cb = kw["stream_callback"] cb = kw["stream_callback"]
cb("r") tid = f"task-{next(counter)}"
cb(STOP) task_ids.append(tid)
cb([(tid, "r"), (tid, STOP)])
return tid
with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
instance = MockSched.return_value
instance.add_task.side_effect = fake_add instance.add_task.side_effect = fake_add
instance.remove_task.return_value = [] instance.remove_task.return_value = []
eng = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=2) eng = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=2)
results = eng.generate(["hello", "world"]) results = eng.generate(["hello", "world"])
assert results == ["r", "r"] assert results == ["r", "r"]
assert task_ids == ["task-0", "task-1"]
def test_engine_generate_zero_max_tokens_returns_empty(): def test_engine_generate_zero_max_tokens_returns_empty():
@@ -294,7 +334,7 @@ def test_generate_captures_calling_backend_context():
def fake_add(prompt, **kwargs): def fake_add(prompt, **kwargs):
captured.append(kwargs["backend"]) captured.append(kwargs["backend"])
kwargs["stream_callback"](STOP) kwargs["stream_callback"]([("task", STOP)])
return "task" return "task"
instance.add_task.side_effect = fake_add instance.add_task.side_effect = fake_add
@@ -328,9 +368,12 @@ def test_build_engine_passes_engine_kwargs_through():
model, _ = make_model("cpu", max_position_embeddings=64) model, _ = make_model("cpu", max_position_embeddings=64)
backend = TorchNativeBackend() backend = TorchNativeBackend()
with patch("astrai.inference.engine.InferenceScheduler") as MockSched: with patch("astrai.inference.engine.InferenceScheduler") as MockSched:
MockSched.return_value.add_task.side_effect = lambda *args, **k: (
k["stream_callback"](STOP) or "task" def fake_add(*args, **k):
) k["stream_callback"]([("task", STOP)])
return "task"
MockSched.return_value.add_task.side_effect = fake_add
engine = build_engine( engine = build_engine(
model=model, model=model,
tokenizer=FakeTokenizer(), tokenizer=FakeTokenizer(),
+60 -1
View File
@@ -4,7 +4,23 @@ from unittest.mock import MagicMock
import pytest import pytest
from astrai.inference import STOP, Task, TaskManager, TaskStatus from astrai.inference import (
STOP,
BatchedStreamCallback,
Task,
TaskManager,
TaskStatus,
)
class RecordingSink(BatchedStreamCallback):
"""Batch-aware callback capturing every dispatch as one batch."""
def __init__(self):
self.batches = []
def __call__(self, events):
self.batches.append(events)
def _make_mock_tokenizer(): def _make_mock_tokenizer():
@@ -217,3 +233,46 @@ def test_task_manager_cancel_active_task_delivers_stop_callback():
immediate, cancelled = tm.cancel_task(task_id) immediate, cancelled = tm.cancel_task(task_id)
assert cancelled and immediate == [] assert cancelled and immediate == []
assert received == [STOP] assert received == [STOP]
def test_invoke_callbacks_batches_sink_events_and_keeps_plain_per_token():
tm = TaskManager(tokenizer=_make_mock_tokenizer())
plain = []
tid_plain = tm.add_task("plain", stream_callback=plain.append)
sink = RecordingSink()
tid_a = tm.add_task("sink a", stream_callback=sink)
tid_b = tm.add_task("sink b", stream_callback=sink)
tm.invoke_callbacks(
[
(tid_a, "x"),
(tid_plain, "p"),
(tid_b, "y"),
("unknown-task", "dropped"),
(tid_a, STOP),
]
)
assert plain == ["p"]
assert sink.batches == [[(tid_a, "x"), (tid_b, "y"), (tid_a, STOP)]]
def test_invoke_callback_delivers_single_event_to_batched_sink():
tm = TaskManager(tokenizer=_make_mock_tokenizer())
sink = RecordingSink()
task_id = tm.add_task("test", stream_callback=sink)
tm.invoke_callback(task_id, STOP)
assert sink.batches == [[(task_id, STOP)]]
def test_cancel_delivers_batched_stop_to_sink():
tm = TaskManager(tokenizer=_make_mock_tokenizer())
sink = RecordingSink()
task_id = tm.add_task("test", stream_callback=sink)
immediate, cancelled = tm.cancel_task(task_id)
assert cancelled
assert sink.batches == [[(task_id, STOP)]]