[Scheduler] Unify per-iteration request intake into ingest_requests() (#38389)

This commit is contained in:
Liangsheng Yin
2026-09-07 20:02:02 -07:00
committed by GitHub
parent 2bf04f3a67
commit b23d835048
12 changed files with 44 additions and 56 deletions
@@ -1154,8 +1154,7 @@ class TestMlxOverlapScheduler(unittest.TestCase):
raise _StopLoop
scheduler = SchedulerMlxOverlapMixin.__new__(SchedulerMlxOverlapMixin)
scheduler.request_receiver = SimpleNamespace(recv_requests=lambda: [])
scheduler.process_input_requests = lambda recv_reqs: None
scheduler.ingest_requests = lambda: []
scheduler.gracefully_exit = False
scheduler._engine_paused = False
scheduler.forward_ct = 0
@@ -118,7 +118,7 @@ class TestOverlapLoopStampsLaunchTs(unittest.TestCase):
scheduler._engine_paused = False
scheduler.waiting_queue = []
scheduler.result_queue = deque()
scheduler.request_receiver.recv_requests.side_effect = recv_side_effect
scheduler.ingest_requests.side_effect = recv_side_effect
result = MagicMock()
result.next_token_ids = None
scheduler.tp_worker.finalize_mlx_result.return_value = result
@@ -269,7 +269,7 @@ class TestOverlapLoopGracefulExit(unittest.TestCase):
scheduler._engine_paused = False
scheduler.waiting_queue = []
scheduler.result_queue = deque()
scheduler.request_receiver.recv_requests.side_effect = recv_side_effect
scheduler.ingest_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.
@@ -296,7 +296,7 @@ class TestOverlapLoopGracefulExit(unittest.TestCase):
) as synchronize:
SchedulerMlxOverlapMixin.event_loop_overlap_mlx(scheduler)
self.assertEqual(scheduler.request_receiver.recv_requests.call_count, 1)
self.assertEqual(scheduler.ingest_requests.call_count, 1)
synchronize.assert_called_once_with()
def test_loop_exits_when_shutdown_arrives_while_paused(self):
@@ -317,7 +317,7 @@ class TestOverlapLoopGracefulExit(unittest.TestCase):
) as synchronize:
SchedulerMlxOverlapMixin.event_loop_overlap_mlx(scheduler)
self.assertEqual(scheduler.request_receiver.recv_requests.call_count, 1)
self.assertEqual(scheduler.ingest_requests.call_count, 1)
scheduler.get_next_batch_to_run.assert_not_called()
synchronize.assert_called_once_with()
@@ -96,12 +96,9 @@ class TestMambaBoundaryMaskReuse(unittest.TestCase):
scheduler = Scheduler.__new__(Scheduler)
scheduler.gracefully_exit = False
scheduler.request_receiver = MagicMock()
scheduler.request_receiver.recv_requests.side_effect = [
[],
[],
StopIteration,
]
scheduler.ingest_requests = MagicMock(
side_effect=[[], [], StopIteration]
)
scheduler.process_input_requests = MagicMock()
scheduler._engine_paused = False
scheduler.running_batch = batch
@@ -121,7 +121,6 @@ def _receiver(tp_size: int = 1) -> SchedulerRequestReceiver:
max_recv_per_poll=-1,
stream_output=lambda *args, **kwargs: None,
get_last_batch=lambda: None,
poll_timeout_aborts=lambda: [],
)
@@ -63,7 +63,6 @@ def _make_receiver(ps: ParallelState) -> SchedulerRequestReceiver:
max_recv_per_poll=-1,
stream_output=lambda *args, **kwargs: None,
get_last_batch=lambda: None,
poll_timeout_aborts=lambda: [],
)