[MLX] Fix overlap-loop request bookkeeping and graceful shutdown (#32447)
Co-authored-by: xiaolin2004 <uwowmhdjwpwpwdhwkw@gmail.com> Co-authored-by: Claude Fable 5 <noreply@anthropic.com> Co-authored-by: R0CKSTAR <yeahdongcn@gmail.com>
This commit is contained in:
co-authored by
xiaolin2004
Claude Fable 5
R0CKSTAR
parent
580b1acbe6
commit
339bef7fad
@@ -16,6 +16,7 @@ the GPU runs both steps back-to-back with no idle gap.
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import time
|
||||
from dataclasses import dataclass
|
||||
from typing import TYPE_CHECKING, List, Optional
|
||||
|
||||
@@ -85,17 +86,18 @@ class MlxPendingJob:
|
||||
class SchedulerMlxOverlapMixin:
|
||||
"""Mixin that adds MLX overlap scheduling to :class:`Scheduler`."""
|
||||
|
||||
def _finalize_mlx_pending_job(self: Scheduler, pending: MlxPendingJob):
|
||||
# Account for this completed forward step. The standard scheduler does
|
||||
# this inside run_batch(), but the MLX overlap loop bypasses run_batch,
|
||||
# so without this forward_ct never advances on MLX. That stalls the
|
||||
# watchdog liveness counter and, more importantly, breaks step-bounded
|
||||
# profiling: _profile_batch_predicate auto-starts/stops based on
|
||||
# forward_ct, so `--profile-steps` (and the server /start_profile
|
||||
# num_steps path) only takes effect once the counter moves here.
|
||||
def _prepare_mlx_launch(self: Scheduler, batch: ScheduleBatch):
|
||||
"""Stamp scheduler bookkeeping before an MLX forward is launched."""
|
||||
# Match run_batch's launch boundary. In particular, the profiler
|
||||
# predicate must run before graph construction / mx.async_eval; running
|
||||
# it while finalizing the previous step profiles at least one queued
|
||||
# decode beyond the requested step count.
|
||||
self.forward_ct += 1
|
||||
self.profiler_manager._profile_batch_predicate(pending.schedule_batch)
|
||||
batch.forward_iter = self.forward_ct
|
||||
batch.launch_ts = time.monotonic()
|
||||
self.profiler_manager._profile_batch_predicate(batch)
|
||||
|
||||
def _finalize_mlx_pending_job(self: Scheduler, pending: MlxPendingJob):
|
||||
result = self.tp_worker.finalize_mlx_result(
|
||||
pending.prefills,
|
||||
pending.extends,
|
||||
@@ -153,6 +155,7 @@ class SchedulerMlxOverlapMixin:
|
||||
pending_next: Optional[MlxPendingJob] = None
|
||||
|
||||
def _launch_fresh(batch: ScheduleBatch) -> MlxPendingJob:
|
||||
self._prepare_mlx_launch(batch)
|
||||
# Materialize batch.input_ids from CPU staging (prefill) or the
|
||||
# FutureMap relay (decode) before the forward. With deferred input
|
||||
# materialization, get_next_batch_to_run leaves input_ids unset; the
|
||||
@@ -160,6 +163,11 @@ class SchedulerMlxOverlapMixin:
|
||||
# loop must do it too, otherwise async_forward_batch_generation_mlx
|
||||
# dereferences a None input_ids.
|
||||
resolve_forward_inputs(batch, self.future_map)
|
||||
# run_batch stamps launch_ts on every scheduler-built forward; the
|
||||
# MLX overlap loop bypasses run_batch, and process_batch_result ->
|
||||
# _record_step_counters subtracts launch_ts unconditionally for
|
||||
# prefill/decode batches. ScheduleBatch.copy() below carries the
|
||||
# stamp to process_batch_result.
|
||||
lazy_tokens, prefills, extends, decode, mode = (
|
||||
self.tp_worker.async_forward_batch_generation_mlx(batch)
|
||||
)
|
||||
@@ -176,24 +184,37 @@ class SchedulerMlxOverlapMixin:
|
||||
|
||||
def _launch_chained(prev: MlxPendingJob) -> MlxPendingJob:
|
||||
assert prev.decode is not None
|
||||
lazy_tokens, prefills, extends, decode, mode = (
|
||||
self.tp_worker.async_chained_decode_mlx(prev.decode)
|
||||
)
|
||||
# Composition is identical to prev: reuse a fresh batch copy
|
||||
# of the same underlying ScheduleBatch so process_batch_result
|
||||
# updates the same req objects with the new token.
|
||||
batch_copy = prev.batch_copy.copy()
|
||||
self._prepare_mlx_launch(batch_copy)
|
||||
# Keep the live scheduler batch's iteration aligned: when the
|
||||
# chain breaks, prepare_for_decode() may run SWA maintenance
|
||||
# before the next fresh launch gets a chance to re-stamp it.
|
||||
prev.schedule_batch.forward_iter = batch_copy.forward_iter
|
||||
lazy_tokens, prefills, extends, decode, mode = (
|
||||
self.tp_worker.async_chained_decode_mlx(prev.decode)
|
||||
)
|
||||
return MlxPendingJob(
|
||||
lazy_tokens=lazy_tokens,
|
||||
prefills=prefills,
|
||||
extends=extends,
|
||||
decode=decode,
|
||||
mode=mode,
|
||||
batch_copy=prev.batch_copy.copy(),
|
||||
batch_copy=batch_copy,
|
||||
schedule_batch=prev.schedule_batch,
|
||||
reqs=prev.reqs,
|
||||
)
|
||||
|
||||
while True:
|
||||
if self.gracefully_exit:
|
||||
# A lookahead job may already be queued by mx.async_eval but
|
||||
# not finalized. Drain Metal work before the scheduler starts
|
||||
# releasing host resources during graceful teardown.
|
||||
mx.synchronize()
|
||||
break
|
||||
|
||||
recv_reqs = self.request_receiver.recv_requests()
|
||||
self.process_input_requests(recv_reqs)
|
||||
if self._engine_paused:
|
||||
|
||||
Reference in New Issue
Block a user