diff --git a/astrai/inference/__init__.py b/astrai/inference/__init__.py index 5e3edc6..c7bec1f 100644 --- a/astrai/inference/__init__.py +++ b/astrai/inference/__init__.py @@ -17,12 +17,20 @@ from astrai.inference.network import get_app, run_server from astrai.inference.runtime.executor import Executor from astrai.inference.runtime.sample import sample 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__ = [ "InferenceEngine", "build_engine", "InferenceScheduler", + "BatchedStreamCallback", "GenerationResult", "Executor", "STOP", diff --git a/astrai/inference/engine.py b/astrai/inference/engine.py index 17d2d7e..4cd70e5 100644 --- a/astrai/inference/engine.py +++ b/astrai/inference/engine.py @@ -13,7 +13,7 @@ import torch.nn as nn from astrai.extension import ATTN_BACKEND, AttentionBackend, get_backend from astrai.inference.cache import PagePool 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.tokenize import AutoTokenizer @@ -33,15 +33,27 @@ class GenerateResult: self._total = count 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: - self.tokens.append((idx, token)) - if token is not STOP: - self.results[idx] += token - else: - if not self._done[idx]: - self._done[idx] = True - self._completed += 1 - self._cond.notify_all() + for idx, token in items: + self.tokens.append((idx, token)) + if token is STOP: + if not self._done[idx]: + self._done[idx] = True + self._completed += 1 + self._cond.notify_all() + else: + self.results[idx] += token self._event.set() def pop_all(self) -> List[Tuple[int, str]]: @@ -69,6 +81,47 @@ class GenerateResult: 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: """Unified inference engine backed by continuous-batching scheduler.""" @@ -147,6 +200,7 @@ class InferenceEngine: ) -> AsyncGenerator[str, None]: request_backend = get_backend(use_default=False) result = GenerateResult() + sink = _ResultSink(result) task_id = self.scheduler.add_task( prompt=prompt, max_tokens=max_tokens, @@ -156,8 +210,9 @@ class InferenceEngine: frequency_penalty=frequency_penalty, rep_window=rep_window, backend=request_backend, - stream_callback=result.append, + stream_callback=sink, ) + sink.bind(task_id, 0) async def _agen(): finished = False @@ -191,8 +246,10 @@ class InferenceEngine: n = len(prompts) request_backend = get_backend(use_default=False) result = GenerateResult(count=n) - task_ids = [ - self.scheduler.add_task( + sink = _ResultSink(result) + task_ids = [] + for i, p in enumerate(prompts): + task_id = self.scheduler.add_task( prompt=p, max_tokens=max_tokens, temperature=temperature, @@ -201,10 +258,10 @@ class InferenceEngine: frequency_penalty=frequency_penalty, rep_window=rep_window, 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: try: diff --git a/astrai/inference/scheduler.py b/astrai/inference/scheduler.py index f22d8bc..0070cf2 100644 --- a/astrai/inference/scheduler.py +++ b/astrai/inference/scheduler.py @@ -264,17 +264,19 @@ class InferenceScheduler: decoded, aborted = self._stepper.step(active) - for t in aborted: - self._task_mgr.invoke_callback(t.task_id, STOP) - + # One dispatch per step: batch-aware sinks take their + # 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: if t.status == TaskStatus.ABORTED: continue new_text = t.decode_new_token(self._task_mgr.tokenizer) 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): - 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: self._stop_event.set() diff --git a/astrai/inference/task.py b/astrai/inference/task.py index 1156563..00b1db3 100644 --- a/astrai/inference/task.py +++ b/astrai/inference/task.py @@ -1,6 +1,7 @@ import threading import time import uuid +from abc import ABC, abstractmethod from collections import deque from dataclasses import dataclass from enum import Enum @@ -145,6 +146,22 @@ class Task: 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: """Thread-safe task queues and lifecycle transitions (no page ops).""" @@ -256,7 +273,10 @@ class TaskManager: immediate = [task] if cancelled and callback is not None: - callback(STOP) + if isinstance(callback, BatchedStreamCallback): + callback([(task_id, STOP)]) + else: + callback(STOP) return immediate, cancelled def remove_task(self, task_id: str) -> List[Task]: @@ -264,10 +284,40 @@ class TaskManager: immediate, _ = self.cancel_task(task_id) return immediate - def invoke_callback(self, task_id: str, token: str): + def invoke_callback(self, task_id: str, token: Any): with self._lock: 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) def get_stats(self) -> Dict[str, Any]: diff --git a/tests/inference/test_engine.py b/tests/inference/test_engine.py index f0c0624..93353c8 100644 --- a/tests/inference/test_engine.py +++ b/tests/inference/test_engine.py @@ -1,6 +1,7 @@ """Unit tests for GenerateResult accumulator and InferenceEngine.generate().""" import asyncio +import itertools import threading from unittest.mock import MagicMock, patch @@ -8,7 +9,12 @@ import pytest from astrai.extension import TorchNativeBackend, attn_backend 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 @@ -50,6 +56,33 @@ def test_result_stop_does_not_double_count(): 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(): r = GenerateResult(count=2) r.append("a", 0) @@ -118,8 +151,8 @@ def test_engine_generate_non_streaming_single(): def fake_add(prompt, **kw): cb = kw["stream_callback"] - cb("response") - cb(STOP) + cb([("task-1", "response"), ("task-1", STOP)]) + return "task-1" instance.add_task.side_effect = fake_add instance.remove_task.return_value = [] @@ -136,6 +169,7 @@ def test_engine_generate_streaming_yields_tokens(): def capture_cb(prompt, **kw): callbacks_saved.append(kw.get("stream_callback")) + return "task-0" with patch("astrai.inference.engine.InferenceScheduler") as MockSched: instance = MockSched.return_value @@ -146,9 +180,9 @@ def test_engine_generate_streaming_yields_tokens(): gen = eng.generate("hello", stream=True) cb = callbacks_saved[0] - cb("t1") - cb("t2") - cb(STOP) + cb([("task-0", "t1")]) + cb([("task-0", "t2")]) + cb([("task-0", STOP)]) tokens = list(gen) 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) stream = engine.generate("hello", stream=True) - callbacks_saved[0]("t1") + callbacks_saved[0]([("task-1", "t1")]) assert next(stream) == "t1" stream.close() @@ -182,6 +216,7 @@ def test_engine_generate_async_yields_tokens_until_stop(): def capture_cb(prompt, **kw): callbacks_saved.append(kw.get("stream_callback")) + return "task-0" with patch("astrai.inference.engine.InferenceScheduler") as MockSched: instance = MockSched.return_value @@ -198,9 +233,9 @@ def test_engine_generate_async_yields_tokens_until_stop(): return out cb = callbacks_saved[0] - cb("t1") - cb("t2") - cb(STOP) + cb([("task-0", "t1")]) + cb([("task-0", "t2")]) + cb([("task-0", STOP)]) 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) stream = engine.generate_async("hello") - callbacks_saved[0]("t1") + callbacks_saved[0]([("task-1", "t1")]) async def consume_then_close(): assert await anext(stream) == "t1" @@ -233,20 +268,25 @@ def test_engine_async_close_cancels_unfinished_task(): def test_engine_generate_non_streaming_batch(): mock_model, mock_tokenizer = _make_engine_mocks(decode="r") + counter = itertools.count() + task_ids = [] + + def fake_add(prompt, **kw): + cb = kw["stream_callback"] + tid = f"task-{next(counter)}" + task_ids.append(tid) + cb([(tid, "r"), (tid, STOP)]) + return tid + with patch("astrai.inference.engine.InferenceScheduler") as MockSched: instance = MockSched.return_value - - def fake_add(prompt, **kw): - cb = kw["stream_callback"] - cb("r") - cb(STOP) - instance.add_task.side_effect = fake_add instance.remove_task.return_value = [] eng = InferenceEngine(mock_model, mock_tokenizer, max_batch_size=2) results = eng.generate(["hello", "world"]) assert results == ["r", "r"] + assert task_ids == ["task-0", "task-1"] 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): captured.append(kwargs["backend"]) - kwargs["stream_callback"](STOP) + kwargs["stream_callback"]([("task", STOP)]) return "task" 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) backend = TorchNativeBackend() 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( model=model, tokenizer=FakeTokenizer(), diff --git a/tests/inference/test_task.py b/tests/inference/test_task.py index 97ab2a4..02ecf39 100644 --- a/tests/inference/test_task.py +++ b/tests/inference/test_task.py @@ -4,7 +4,23 @@ from unittest.mock import MagicMock 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(): @@ -217,3 +233,46 @@ def test_task_manager_cancel_active_task_delivers_stop_callback(): immediate, cancelled = tm.cancel_task(task_id) assert cancelled and immediate == [] 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)]]