From 7b879593bbff72cee1f536a47ab2601255ed563c Mon Sep 17 00:00:00 2001 From: Alex Cheema Date: Fri, 13 Feb 2026 07:29:13 -0800 Subject: [PATCH] fix: update continuous batching types after main merge Replace ChatCompletionTaskParams with TextGenerationTaskParams and ChatCompletion with TextGeneration to match the refactored type hierarchy from main. Add missing usage parameter to GenerationResponse constructors and add type annotations to StreamingDetokenizer stubs. Co-Authored-By: Claude Opus 4.6 --- .mlx_typings/mlx_lm/tokenizer_utils.pyi | 8 ++--- conftest.py | 1 + .../engines/mlx/generator/batch_engine.py | 17 ++++++---- .../test_runner/test_continuous_batching.py | 31 ++++++++++--------- 4 files changed, 33 insertions(+), 24 deletions(-) create mode 100644 conftest.py diff --git a/.mlx_typings/mlx_lm/tokenizer_utils.pyi b/.mlx_typings/mlx_lm/tokenizer_utils.pyi index dd12af3b..1326f059 100644 --- a/.mlx_typings/mlx_lm/tokenizer_utils.pyi +++ b/.mlx_typings/mlx_lm/tokenizer_utils.pyi @@ -39,11 +39,11 @@ class StreamingDetokenizer: """ __slots__ = ... - def reset(self): ... - def add_token(self, token): ... - def finalize(self): ... + def reset(self) -> None: ... + def add_token(self, token: int) -> None: ... + def finalize(self) -> None: ... @property - def last_segment(self): + def last_segment(self) -> str: """Return the last segment of readable text since last time this property was accessed.""" class NaiveStreamingDetokenizer(StreamingDetokenizer): diff --git a/conftest.py b/conftest.py new file mode 100644 index 00000000..c6a5375d --- /dev/null +++ b/conftest.py @@ -0,0 +1 @@ +collect_ignore = ["tests/start_distributed_test.py"] diff --git a/src/exo/worker/engines/mlx/generator/batch_engine.py b/src/exo/worker/engines/mlx/generator/batch_engine.py index d766569e..09bbfd74 100644 --- a/src/exo/worker/engines/mlx/generator/batch_engine.py +++ b/src/exo/worker/engines/mlx/generator/batch_engine.py @@ -11,7 +11,8 @@ from mlx_lm.tokenizer_utils import StreamingDetokenizer, TokenizerWrapper from exo.shared.types.api import FinishReason, GenerationStats from exo.shared.types.common import CommandId from exo.shared.types.memory import Memory -from exo.shared.types.tasks import ChatCompletionTaskParams, TaskId +from exo.shared.types.tasks import TaskId +from exo.shared.types.text_generation import TextGenerationTaskParams from exo.shared.types.worker.runner_response import GenerationResponse from exo.worker.engines.mlx import Model from exo.worker.engines.mlx.constants import MAX_TOKENS @@ -60,7 +61,7 @@ class BatchGenerationEngine: self.max_tokens = max_tokens self.active_requests: dict[int, ActiveRequest] = {} self._pending_inserts: list[ - tuple[CommandId, TaskId, ChatCompletionTaskParams] + tuple[CommandId, TaskId, TextGenerationTaskParams] ] = [] self._pending_completions: list[ int @@ -93,7 +94,7 @@ class BatchGenerationEngine: self, command_id: CommandId, task_id: TaskId, - task_params: ChatCompletionTaskParams, + task_params: TextGenerationTaskParams, ) -> None: """Queue a request for insertion. Only rank 0 should call this. @@ -117,7 +118,7 @@ class BatchGenerationEngine: Batches all pending inserts into a single batch_gen.insert() call for efficient prefill batching. """ - inserts_to_process: list[tuple[CommandId, TaskId, ChatCompletionTaskParams]] + inserts_to_process: list[tuple[CommandId, TaskId, TextGenerationTaskParams]] if not self.is_distributed: # Non-distributed: just insert directly from pending @@ -149,7 +150,7 @@ class BatchGenerationEngine: tokens: list[int] = self.tokenizer.encode( prompt_str, add_special_tokens=False ) - max_tokens = params.max_tokens or self.max_tokens + max_tokens = params.max_output_tokens or self.max_tokens all_tokens.append(tokens) all_max_tokens.append(max_tokens) @@ -237,7 +238,11 @@ class BatchGenerationEngine: command_id=req.command_id, task_id=req.task_id, response=GenerationResponse( - text=text, token=token, finish_reason=finish_reason, stats=stats + text=text, + token=token, + finish_reason=finish_reason, + stats=stats, + usage=None, ), ) ) diff --git a/src/exo/worker/tests/unittests/test_runner/test_continuous_batching.py b/src/exo/worker/tests/unittests/test_runner/test_continuous_batching.py index 0887d29d..dbb22005 100644 --- a/src/exo/worker/tests/unittests/test_runner/test_continuous_batching.py +++ b/src/exo/worker/tests/unittests/test_runner/test_continuous_batching.py @@ -11,6 +11,7 @@ NOTE: These tests require the continuous-batching runner architecture (BatchGenerationEngine) which is not yet integrated with main. """ +# ruff: noqa: E402 # pyright: reportAny=false # pyright: reportUnknownArgumentType=false # pyright: reportUnknownMemberType=false @@ -20,7 +21,7 @@ NOTE: These tests require the continuous-batching runner architecture import pytest pytest.skip( - "continuous batching runner not yet updated for main branch types", + "continuous batching runner not yet integrated with main branch runner", allow_module_level=True, ) @@ -28,7 +29,6 @@ from typing import Any from unittest.mock import MagicMock import exo.worker.runner.runner as mlx_runner -from exo.shared.types.api import ChatCompletionMessage from exo.shared.types.common import CommandId, NodeId from exo.shared.types.events import ( Event, @@ -36,8 +36,6 @@ from exo.shared.types.events import ( TaskStatusUpdated, ) from exo.shared.types.tasks import ( - ChatCompletion, - ChatCompletionTaskParams, ConnectToGroup, LoadModel, Shutdown, @@ -45,7 +43,9 @@ from exo.shared.types.tasks import ( Task, TaskId, TaskStatus, + TextGeneration, ) +from exo.shared.types.text_generation import InputMessage, TextGenerationTaskParams from exo.shared.types.worker.runner_response import GenerationResponse from exo.shared.types.worker.runners import RunnerRunning from exo.utils.channels import mp_channel @@ -74,7 +74,7 @@ class FakeBatchEngineWithTokens: def __init__(self, *_args: Any, **_kwargs: Any): self._active_requests: dict[int, tuple[CommandId, TaskId, int, int]] = {} self._pending_inserts: list[ - tuple[CommandId, TaskId, ChatCompletionTaskParams] + tuple[CommandId, TaskId, TextGenerationTaskParams] ] = [] self._uid_counter = 0 self._tokens_per_request = 3 # Default: generate 3 tokens before completing @@ -84,7 +84,7 @@ class FakeBatchEngineWithTokens: self, command_id: CommandId, task_id: TaskId, - task_params: ChatCompletionTaskParams, + task_params: TextGenerationTaskParams, ) -> None: """Queue a request for insertion.""" self._pending_inserts.append((command_id, task_id, task_params)) @@ -106,12 +106,14 @@ class FakeBatchEngineWithTokens: self, command_id: CommandId, task_id: TaskId, - task_params: ChatCompletionTaskParams | None, + task_params: TextGenerationTaskParams | None, ) -> int: uid = self._uid_counter self._uid_counter += 1 # Track: (command_id, task_id, tokens_generated, max_tokens) - max_tokens = task_params.max_tokens if task_params else self._tokens_per_request + max_tokens = ( + task_params.max_output_tokens if task_params else self._tokens_per_request + ) self._active_requests[uid] = (command_id, task_id, 0, max_tokens or 3) return uid @@ -144,6 +146,7 @@ class FakeBatchEngineWithTokens: token=tokens_gen, text=text, finish_reason=finish_reason, + usage=None, ), ) ) @@ -243,15 +246,15 @@ WARMUP_TASK = StartWarmup(task_id=TaskId("warmup"), instance_id=INSTANCE_1_ID) def make_chat_task( task_id: str, command_id: str, max_tokens: int = 3 -) -> ChatCompletion: - return ChatCompletion( +) -> TextGeneration: + return TextGeneration( task_id=TaskId(task_id), command_id=CommandId(command_id), - task_params=ChatCompletionTaskParams( - model=str(MODEL_A_ID), - messages=[ChatCompletionMessage(role="user", content="hello")], + task_params=TextGenerationTaskParams( + model=MODEL_A_ID, + input=[InputMessage(role="user", content="hello")], stream=True, - max_tokens=max_tokens, + max_output_tokens=max_tokens, ), instance_id=INSTANCE_1_ID, )