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:
@@ -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
@@ -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:
|
||||||
|
|||||||
@@ -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()
|
||||||
|
|||||||
@@ -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]:
|
||||||
|
|||||||
@@ -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(),
|
||||||
|
|||||||
@@ -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)]]
|
||||||
|
|||||||
Reference in New Issue
Block a user