Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
57 changes: 54 additions & 3 deletions areno/api/backend/mlx/generation.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,12 +39,14 @@ class _Request:
features: list[dict | None] | None
future: Future[list[RolloutResult]] = field(default_factory=Future)
handles: list[tuple[object, int]] = field(default_factory=list)
handle_to_prompt_idx: dict[tuple[object, int], int] = field(default_factory=dict)
tokens: dict[tuple[object, int], list[int]] = field(default_factory=dict)
logprobs: dict[tuple[object, int], list[float]] = field(default_factory=dict)
finished: set[tuple[object, int]] = field(default_factory=set)
expanded_prompts: list[list[int]] = field(default_factory=list)
expanded_features: list[dict | None] = field(default_factory=list)
next_insert: int = 0
stream_queue: queue.Queue | None = None


@dataclass(slots=True)
Expand Down Expand Up @@ -231,6 +233,38 @@ async def submit_async(
) -> list[RolloutResult]:
return await asyncio.wrap_future(self.submit(prompt_tokens, n_samples, sampling_params, prompt_features))

def submit_stream(
self,
prompt_tokens: list[list[int]],
n_samples: int,
sampling_params: SamplingParams,
prompt_features: list[dict | None] | None = None,
) -> queue.Queue:
"""Submit a streaming rollout request.

Returns a thread-safe :class:`queue.Queue` that receives
``(prompt_idx, token_id, finish_reason)`` tuples as tokens are
generated by the decode thread, followed by a ``None`` sentinel
when all sequences complete. On scheduler failure, an exception
object is placed on the queue before the sentinel.
"""
if self._closed:
raise RuntimeError("MLX rollout scheduler is closed")
if self._failure is not None:
raise RuntimeError("MLX rollout scheduler failed") from self._failure
if n_samples < 1:
raise ValueError("n_samples must be positive")
if prompt_features is not None and len(prompt_features) != len(prompt_tokens):
raise ValueError("prompt_features must align with prompt_tokens")
if not prompt_tokens:
q: queue.Queue = queue.Queue()
q.put(None)
return q
request = _Request(prompt_tokens, n_samples, sampling_params, prompt_features)
request.stream_queue = queue.Queue()
self._commands.put(request)
return request.stream_queue

def drop_state(self) -> None:
"""Release completed KV and allocator caches without replacing the scheduler."""

Expand Down Expand Up @@ -369,21 +403,35 @@ def _insert(self, request: _Request, start: int, end: int) -> None:
request.handles.extend(handles)
request.tokens.update((handle, []) for handle in handles)
request.logprobs.update((handle, []) for handle in handles)
for handle in handles:
for i, handle in enumerate(handles):
self._requests_by_handle[handle] = request
if request.stream_queue is not None:
request.handle_to_prompt_idx[handle] = (start + i) // request.n_samples

def _record_response(self, key: object, generator: Any, response: Any) -> None:
handle = (key, int(response.uid))
request = self._requests_by_handle[handle]

is_finished = response.finish_reason is not None

if response.finish_reason != "stop":
request.tokens[handle].append(int(response.token))
token_id = int(response.token)
request.tokens[handle].append(token_id)
request.logprobs[handle].append(generator.token_logprob(response))
if response.finish_reason is None:
if request.stream_queue is not None:
prompt_idx = request.handle_to_prompt_idx.get(handle, 0)
stream_fr: str | None = response.finish_reason if is_finished else None
request.stream_queue.put((prompt_idx, token_id, stream_fr))

if not is_finished:
return

request.finished.add(handle)
self._requests_by_handle.pop(handle, None)
if len(request.finished) == len(request.expanded_prompts):
request.future.set_result(_request_result(request))
if request.stream_queue is not None:
request.stream_queue.put(None)

def _record_decode_progress(self, token_delta: int) -> None:
"""Emit the same throttled decode progress line as the CUDA backend."""
Expand Down Expand Up @@ -428,6 +476,9 @@ def _fail_all(self, exc: BaseException) -> None:
for request in requests.values():
if not request.future.done():
request.future.set_exception(exc)
if request.stream_queue is not None:
request.stream_queue.put(exc)
request.stream_queue.put(None)


def _request_result(request: _Request) -> list[RolloutResult]:
Expand Down
Loading