From 339bef7fad7585e57a73d647ba089122e1f7e4a3 Mon Sep 17 00:00:00 2001 From: Donkey Kane Date: Wed, 29 Jul 2026 10:04:53 +0800 Subject: [PATCH] [MLX] Fix overlap-loop request bookkeeping and graceful shutdown (#32447) Co-authored-by: xiaolin2004 Co-authored-by: Claude Fable 5 Co-authored-by: R0CKSTAR --- .../hardware_backend/mlx/scheduler_mixin.py | 47 ++- .../mlx/test_attention_patching.py | 10 +- .../mlx/test_scheduler_mixin.py | 270 ++++++++++++++++-- 3 files changed, 292 insertions(+), 35 deletions(-) diff --git a/python/sglang/srt/hardware_backend/mlx/scheduler_mixin.py b/python/sglang/srt/hardware_backend/mlx/scheduler_mixin.py index 581f15fb2..c462c0342 100644 --- a/python/sglang/srt/hardware_backend/mlx/scheduler_mixin.py +++ b/python/sglang/srt/hardware_backend/mlx/scheduler_mixin.py @@ -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: diff --git a/test/registered/unit/hardware_backend/mlx/test_attention_patching.py b/test/registered/unit/hardware_backend/mlx/test_attention_patching.py index db92d68fe..36d874d02 100644 --- a/test/registered/unit/hardware_backend/mlx/test_attention_patching.py +++ b/test/registered/unit/hardware_backend/mlx/test_attention_patching.py @@ -1111,7 +1111,6 @@ class TestMlxOverlapScheduler(unittest.TestCase): self.assertTrue(torch.equal(schedule_batch.input_ids, token_ids)) self.assertIs(scheduler.processed_batch, batch_copy) self.assertIs(scheduler.processed_result, scheduler.tp_worker.result) - self.assertEqual(scheduler.forward_ct, 1) def test_overlap_loop_materializes_prefill_input_ids(self): # Regression: the MLX overlap loop must materialize batch.input_ids @@ -1132,7 +1131,12 @@ class TestMlxOverlapScheduler(unittest.TestCase): scheduler = SchedulerMlxOverlapMixin.__new__(SchedulerMlxOverlapMixin) scheduler.request_receiver = SimpleNamespace(recv_requests=lambda: []) scheduler.process_input_requests = lambda recv_reqs: None + scheduler.gracefully_exit = False scheduler._engine_paused = False + scheduler.forward_ct = 0 + scheduler.profiler_manager = SimpleNamespace( + _profile_batch_predicate=lambda batch: None + ) scheduler.waiting_queue = [] scheduler.result_queue = deque() scheduler.future_map = SimpleNamespace() @@ -1504,9 +1508,7 @@ if _HAS_MLX: self.last_batch = None self.processed_batch = None self.processed_result = None - # _finalize_mlx_pending_job now advances forward_ct and runs the - # profiler batch predicate (mirroring run_batch); stub both so the - # overlap accounting added in #29217 has something to call. + # Launch bookkeeping mirrors run_batch before each MLX forward. self.forward_ct = 0 self.profiler_manager = SimpleNamespace( _profile_batch_predicate=lambda batch: None diff --git a/test/registered/unit/hardware_backend/mlx/test_scheduler_mixin.py b/test/registered/unit/hardware_backend/mlx/test_scheduler_mixin.py index cc620f876..4fbbd5912 100644 --- a/test/registered/unit/hardware_backend/mlx/test_scheduler_mixin.py +++ b/test/registered/unit/hardware_backend/mlx/test_scheduler_mixin.py @@ -1,11 +1,9 @@ """Unit tests for the MLX overlap scheduler mixin (hardware_backend/mlx/scheduler_mixin.py). Covers: - - _finalize_mlx_pending_job advances forward_ct once per completed step - - _finalize_mlx_pending_job calls the profiler batch predicate with the - finalized batch, so step-bounded profiling (``--profile-steps`` / - ``/start_profile`` num_steps) can auto-stop on the MLX overlap loop, which - bypasses the standard Scheduler.run_batch(). + - Every MLX launch advances forward_ct and stamps forward_iter/launch_ts. + - The profiler predicate runs before the async forward is enqueued, matching + Scheduler.run_batch() so step-bounded profiling stops on the right step. Skips on non-Apple-Silicon platforms and when ``mlx`` is missing (importing scheduler_mixin requires ``mlx.core``). @@ -16,7 +14,7 @@ from __future__ import annotations import importlib.util import platform import unittest -from unittest.mock import MagicMock +from unittest.mock import MagicMock, patch from sglang.test.ci.ci_register import register_cpu_ci, register_mlx_ci @@ -29,8 +27,8 @@ _SKIP_REASON = "requires Apple Silicon and mlx" @unittest.skipUnless(_IS_APPLE_SILICON and _HAS_MLX, _SKIP_REASON) -class TestFinalizeMlxPendingJob(unittest.TestCase): - """forward_ct accounting + profiler predicate wiring in the overlap loop.""" +class TestMlxLaunchBookkeeping(unittest.TestCase): + """run_batch-style bookkeeping for the MLX overlap loop.""" def _make_scheduler(self): scheduler = MagicMock() @@ -40,26 +38,24 @@ class TestFinalizeMlxPendingJob(unittest.TestCase): scheduler.tp_worker.finalize_mlx_result.return_value = result return scheduler - def test_finalize_advances_forward_ct_and_runs_predicate(self): + def test_prepare_launch_advances_forward_ct_and_runs_predicate(self): from sglang.srt.hardware_backend.mlx.scheduler_mixin import ( SchedulerMlxOverlapMixin, ) scheduler = self._make_scheduler() - pending = MagicMock() + batch = MagicMock() - SchedulerMlxOverlapMixin._finalize_mlx_pending_job(scheduler, pending) + SchedulerMlxOverlapMixin._prepare_mlx_launch(scheduler, batch) - # Standard run_batch() advances forward_ct and runs the profiler - # predicate; the MLX overlap loop must do the same here. self.assertEqual(scheduler.forward_ct, 1) + self.assertEqual(batch.forward_iter, 1) + self.assertIsInstance(batch.launch_ts, float) scheduler.profiler_manager._profile_batch_predicate.assert_called_once_with( - pending.schedule_batch + batch ) - # The rest of finalization still runs. - scheduler.process_batch_result.assert_called_once() - def test_forward_ct_advances_once_per_step(self): + def test_forward_ct_advances_once_per_launch(self): from sglang.srt.hardware_backend.mlx.scheduler_mixin import ( SchedulerMlxOverlapMixin, ) @@ -67,13 +63,251 @@ class TestFinalizeMlxPendingJob(unittest.TestCase): scheduler = self._make_scheduler() for expected in (1, 2, 3): - SchedulerMlxOverlapMixin._finalize_mlx_pending_job(scheduler, MagicMock()) + SchedulerMlxOverlapMixin._prepare_mlx_launch(scheduler, MagicMock()) self.assertEqual(scheduler.forward_ct, expected) self.assertEqual( scheduler.profiler_manager._profile_batch_predicate.call_count, 3 ) + def test_finalize_does_not_double_count_launch(self): + from sglang.srt.hardware_backend.mlx.scheduler_mixin import ( + SchedulerMlxOverlapMixin, + ) + + scheduler = self._make_scheduler() + pending = MagicMock() + + SchedulerMlxOverlapMixin._prepare_mlx_launch(scheduler, pending.batch_copy) + SchedulerMlxOverlapMixin._finalize_mlx_pending_job(scheduler, pending) + + self.assertEqual(scheduler.forward_ct, 1) + self.assertEqual(pending.batch_copy.forward_iter, 1) + scheduler.process_batch_result.assert_called_once() + + +class _StopLoop(Exception): + """Sentinel to break out of the event loop's ``while True``.""" + + +@unittest.skipUnless(_IS_APPLE_SILICON and _HAS_MLX, _SKIP_REASON) +class TestOverlapLoopStampsLaunchTs(unittest.TestCase): + """Every batch the MLX overlap loop launches must carry ``launch_ts``. + + ``Scheduler.run_batch`` stamps ``batch.launch_ts`` on every forward, and + ``process_batch_result`` -> ``_record_step_counters`` subtracts it + unconditionally for prefill/decode batches. The MLX overlap loop bypasses + ``run_batch``, so if its launch paths skip the stamp, the first real + request's result processing raises ``TypeError: float - NoneType`` and + kills the scheduler (health-check requests are filtered from the counters, + which keeps ``/health_generate`` green while every real request crashes). + """ + + def _make_scheduler(self, *, recv_side_effect): + from collections import deque + + from sglang.srt.hardware_backend.mlx.scheduler_mixin import ( + SchedulerMlxOverlapMixin, + ) + + scheduler = MagicMock() + scheduler.forward_ct = 0 + scheduler._prepare_mlx_launch.side_effect = lambda batch: ( + SchedulerMlxOverlapMixin._prepare_mlx_launch(scheduler, batch) + ) + scheduler.gracefully_exit = False + scheduler._engine_paused = False + scheduler.waiting_queue = [] + scheduler.result_queue = deque() + scheduler.request_receiver.recv_requests.side_effect = recv_side_effect + result = MagicMock() + result.next_token_ids = None + scheduler.tp_worker.finalize_mlx_result.return_value = result + return scheduler + + def test_fresh_launch_stamps_launch_ts_before_input_resolution(self): + from sglang.srt.hardware_backend.mlx.scheduler_mixin import ( + SchedulerMlxOverlapMixin, + ) + + scheduler = self._make_scheduler(recv_side_effect=[[], _StopLoop()]) + + batch = MagicMock() + events = [] + scheduler.profiler_manager._profile_batch_predicate.side_effect = ( + lambda _batch: events.append("profile") + ) + launch_ts_at_copy_time = [] + batch.copy.side_effect = lambda: ( + launch_ts_at_copy_time.append(batch.launch_ts), + MagicMock(), + )[1] + plan = MagicMock() + plan.batch_to_run = batch + scheduler.get_next_batch_to_run.return_value = plan + scheduler.tp_worker.async_forward_batch_generation_mlx.side_effect = ( + lambda _batch: ( + events.append("forward"), + (None, [], [], None, "extend"), + )[1] + ) + + with ( + patch( + "sglang.srt.hardware_backend.mlx.scheduler_mixin.time.monotonic", + side_effect=lambda: (events.append("launch_ts"), 1.0)[1], + ), + patch( + "sglang.srt.hardware_backend.mlx.scheduler_mixin.resolve_forward_inputs", + side_effect=lambda *_args: events.append("resolve_inputs"), + ), + self.assertRaises(_StopLoop), + ): + SchedulerMlxOverlapMixin.event_loop_overlap_mlx(scheduler) + + self.assertEqual(events, ["launch_ts", "profile", "resolve_inputs", "forward"]) + self.assertEqual(len(launch_ts_at_copy_time), 1) + self.assertEqual(launch_ts_at_copy_time[0], 1.0) + + def test_chained_launch_restamps_launch_ts(self): + from sglang.srt.hardware_backend.mlx.scheduler_mixin import ( + SchedulerMlxOverlapMixin, + ) + + # Iteration 1: fresh decode launch. Iteration 2: chain a second + # decode on top of it. Iteration 3: stop. + scheduler = self._make_scheduler(recv_side_effect=[[], [], _StopLoop()]) + + events = [] + req = MagicMock() + req.finished.return_value = False + batch = MagicMock() + batch.reqs = [req] + fresh_copy = MagicMock() + batch.copy.return_value = fresh_copy + chained_copy = MagicMock() + chained_copy.launch_ts = None + fresh_copy.copy.return_value = chained_copy + plan = MagicMock() + plan.batch_to_run = batch + scheduler.get_next_batch_to_run.return_value = plan + + pending_decode = MagicMock() + scheduler.tp_worker.async_forward_batch_generation_mlx.return_value = ( + MagicMock(), + [], + [], + pending_decode, + "decode", + ) + scheduler.tp_worker.async_chained_decode_mlx.side_effect = lambda _decode: ( + events.append("chained_forward"), + (MagicMock(), [], [], MagicMock(), "decode"), + )[1] + + launch_times = iter((1.0, 2.0)) + + def record_launch_ts(): + launch_ts = next(launch_times) + events.append(f"launch_ts:{launch_ts}") + return launch_ts + + with ( + patch( + "sglang.srt.hardware_backend.mlx.scheduler_mixin.time.monotonic", + side_effect=record_launch_ts, + ), + patch( + "sglang.srt.hardware_backend.mlx.scheduler_mixin.resolve_forward_inputs" + ), + self.assertRaises(_StopLoop), + ): + SchedulerMlxOverlapMixin.event_loop_overlap_mlx(scheduler) + + scheduler.tp_worker.async_chained_decode_mlx.assert_called_once() + self.assertLess(events.index("launch_ts:2.0"), events.index("chained_forward")) + self.assertEqual(chained_copy.launch_ts, 2.0) + # The live batch only needs the iteration for SWA maintenance before + # the next fresh launch; per-step timing consumes the batch copy. + self.assertEqual(batch.forward_iter, 2) + self.assertEqual(batch.launch_ts, 1.0) + + +@unittest.skipUnless(_IS_APPLE_SILICON and _HAS_MLX, _SKIP_REASON) +class TestOverlapLoopGracefulExit(unittest.TestCase): + """The MLX overlap loop must honor ``gracefully_exit`` like the standard loops. + + ``handle_shutdown`` (ShutdownReq) only sets ``scheduler.gracefully_exit``; + actual teardown happens after the event loop returns — + ``run_scheduler_process``'s ``finally`` calls ``release_host_resources()`` + only once the loop breaks. ``event_loop_normal`` and ``event_loop_overlap`` + check the flag at the top of every iteration; a loop that never checks it + spins forever, so the TokenizerManager's shutdown path times out after its + 15 s grace period and falls back to ``kill_process_tree`` — host resources + never get their user-space release. + """ + + def _make_scheduler(self, *, recv_side_effect): + from collections import deque + + scheduler = MagicMock() + scheduler.forward_ct = 0 + scheduler.gracefully_exit = False + scheduler._engine_paused = False + scheduler.waiting_queue = [] + scheduler.result_queue = deque() + scheduler.request_receiver.recv_requests.side_effect = recv_side_effect + # Model handle_shutdown: processing a non-empty recv batch (the + # ShutdownReq) flips the flag; the loop must notice at the top of the + # next iteration instead of polling forever. + scheduler.process_input_requests.side_effect = lambda reqs: ( + setattr(scheduler, "gracefully_exit", True) if reqs else None + ) + plan = MagicMock() + plan.batch_to_run = None + scheduler.get_next_batch_to_run.return_value = plan + return scheduler + + def test_loop_exits_after_shutdown_req(self): + from sglang.srt.hardware_backend.mlx.scheduler_mixin import ( + SchedulerMlxOverlapMixin, + ) + + # Iteration 1: recv the ShutdownReq stand-in (flag flips inside + # process_input_requests). Iteration 2 must break before polling + # again; the sentinel raising instead means the loop never exits. + scheduler = self._make_scheduler(recv_side_effect=[[MagicMock()], _StopLoop()]) + + with patch( + "sglang.srt.hardware_backend.mlx.scheduler_mixin.mx.synchronize" + ) as synchronize: + SchedulerMlxOverlapMixin.event_loop_overlap_mlx(scheduler) + + self.assertEqual(scheduler.request_receiver.recv_requests.call_count, 1) + synchronize.assert_called_once_with() + + def test_loop_exits_when_shutdown_arrives_while_paused(self): + from sglang.srt.hardware_backend.mlx.scheduler_mixin import ( + SchedulerMlxOverlapMixin, + ) + + # A paused engine still recvs and processes control requests — that is + # how unpause (and shutdown) arrive — but `continue`s past the rest of + # the body. The flag check must sit above the paused-continue, like in + # event_loop_normal/event_loop_overlap, or shutdown during a pause + # spins forever. + scheduler = self._make_scheduler(recv_side_effect=[[MagicMock()], _StopLoop()]) + scheduler._engine_paused = True + + with patch( + "sglang.srt.hardware_backend.mlx.scheduler_mixin.mx.synchronize" + ) as synchronize: + SchedulerMlxOverlapMixin.event_loop_overlap_mlx(scheduler) + + self.assertEqual(scheduler.request_receiver.recv_requests.call_count, 1) + scheduler.get_next_batch_to_run.assert_not_called() + synchronize.assert_called_once_with() + if __name__ == "__main__": unittest.main()