[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
+2 -4
View File
@@ -2504,8 +2504,7 @@ class SchedulerDisaggregationDecodeMixin:
if not self._engine_paused:
self.disagg_decode_prealloc_queue.prefetch_prefill_dp_rank_queries()
# Receive requests
recv_reqs = self.request_receiver.recv_requests()
self.process_input_requests(recv_reqs)
self.ingest_requests()
if self._engine_paused:
self._record_scheduler_state_for_paused_engine()
continue
@@ -2548,8 +2547,7 @@ class SchedulerDisaggregationDecodeMixin:
if not self._engine_paused:
self.disagg_decode_prealloc_queue.prefetch_prefill_dp_rank_queries()
# Receive requests
recv_reqs = self.request_receiver.recv_requests()
self.process_input_requests(recv_reqs)
self.ingest_requests()
if self._engine_paused:
self._record_scheduler_state_for_paused_engine()
continue
+2 -4
View File
@@ -599,8 +599,7 @@ class SchedulerDisaggregationPrefillMixin:
"""A normal scheduler loop for prefill worker in disaggregation mode."""
while True:
# Receive requests
recv_reqs = self.request_receiver.recv_requests()
self.process_input_requests(recv_reqs)
self.ingest_requests()
if self._engine_paused:
self._record_scheduler_state_for_paused_engine()
continue
@@ -639,8 +638,7 @@ class SchedulerDisaggregationPrefillMixin:
while True:
# Receive requests
recv_reqs = self.request_receiver.recv_requests()
self.process_input_requests(recv_reqs)
self.ingest_requests()
if self._engine_paused:
self._record_scheduler_state_for_paused_engine()
continue
@@ -210,8 +210,7 @@ class SchedulerMlxOverlapMixin:
mx.synchronize()
break
recv_reqs = self.request_receiver.recv_requests()
self.process_input_requests(recv_reqs)
self.ingest_requests()
if self._engine_paused:
self._record_scheduler_state_for_paused_engine()
continue
+21 -9
View File
@@ -1895,10 +1895,7 @@ class Scheduler(
break
# Receive requests
recv_reqs = self.request_receiver.recv_requests()
if recv_reqs:
self.metrics_reporter.record_scheduler_active()
self.process_input_requests(recv_reqs)
self.ingest_requests()
if self._engine_paused:
self._record_scheduler_state_for_paused_engine()
continue
@@ -1942,10 +1939,7 @@ class Scheduler(
break
# Receive requests
recv_reqs = self.request_receiver.recv_requests()
if recv_reqs:
self.metrics_reporter.record_scheduler_active()
self.process_input_requests(recv_reqs)
self.ingest_requests()
if self._engine_paused:
self._record_scheduler_state_for_paused_engine()
continue
@@ -2051,6 +2045,25 @@ class Scheduler(
for prev_batch, prev_result in self.result_queue:
self.batch_result_processor.advance_grammar_fsm(prev_result, prev_batch)
def ingest_requests(self) -> List:
"""Receive, broadcast and dispatch this iteration's external input.
The one place a new per-iteration input source belongs; the return
value exists for the pipeline stages that relay requests onward.
"""
local_reqs = []
if (
self.ps.pp_rank == 0
and self.ps.attn_tp_rank == 0
and self.ps.attn_cp_rank == 0
):
local_reqs = self._poll_timeout_aborts()
recv_reqs = self.request_receiver.recv_requests(local_reqs=local_reqs)
if recv_reqs:
self.metrics_reporter.record_scheduler_active()
self.process_input_requests(recv_reqs)
return recv_reqs
@scheduler_stage_method(SCHEDULER_STAGE_PROCESS_REQUESTS)
def process_input_requests(self, recv_reqs: List):
now = time.monotonic()
@@ -2278,7 +2291,6 @@ class Scheduler(
get_last_batch=lambda: self.last_batch,
scripted_scheduler_hook=self.scripted_scheduler_hook,
scheduler_stage_metrics=self.scheduler_stage_metrics,
poll_timeout_aborts=self._poll_timeout_aborts,
)
def init_dp_attn_adapter(self) -> None:
@@ -79,9 +79,6 @@ class SchedulerRequestReceiver:
get_last_batch: Callable[[], Any]
scripted_scheduler_hook: Optional[ScriptedSchedulerHook] = None
scheduler_stage_metrics: Optional[SchedulerStageMetricsRecorder] = None
# Emits AbortReqs for SGLANG_REQ_WAITING_TIMEOUT / _RUNNING_TIMEOUT;
# runs on the rank that owns the waiting queue.
poll_timeout_aborts: Callable[[], List[AbortReq]]
def recv_limit_reached(self, num_recv_reqs: int) -> bool:
if self.max_recv_per_poll < 0:
@@ -90,9 +87,13 @@ class SchedulerRequestReceiver:
@scheduler_stage_method(SCHEDULER_STAGE_RECV_REQUESTS)
def recv_requests(
self,
self, local_reqs: Optional[List[AbortReq]] = None
) -> List[Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput, Any]]:
"""Receive results at tp_rank = 0 and broadcast it to all other TP ranks."""
"""Receive results at tp_rank = 0 and broadcast it to all other TP ranks.
local_reqs are aborts the caller decided on this rank; they ride the
same broadcast as the pulled requests.
"""
if self.scripted_scheduler_hook is not None:
self.scripted_scheduler_hook.step()
@@ -106,16 +107,6 @@ class SchedulerRequestReceiver:
if self.input_blocker is not None:
recv_reqs = self.input_blocker.handle(recv_reqs)
# Decided once and broadcast, so every rank sharing this waiting queue
# drops the same requests in the same iteration.
local_reqs = []
if (
self.ps.pp_rank == 0
and self.ps.attn_tp_rank == 0
and self.ps.attn_cp_rank == 0
):
local_reqs = self.poll_timeout_aborts()
recv_reqs = self._broadcast_reqs_across_ranks(recv_reqs, local_reqs)
if self.ps.pp_rank == 0:
@@ -92,8 +92,7 @@ class SchedulerPPMixin:
next_first_rank_mb_id = (mb_id + self.ps.pp_size) % self.pp_loop_size
next_mb_id = (mb_id + 1) % self.pp_loop_size
with torch.profiler.record_function("recv_requests"):
recv_reqs = self.request_receiver.recv_requests()
self.process_input_requests(recv_reqs)
recv_reqs = self.ingest_requests()
if not self.pp_group.is_last_rank:
self._pp_commit_comm_work(self.send_req_work)
with torch.profiler.record_function("send_reqs_to_next_stage"):
@@ -232,8 +231,7 @@ class SchedulerPPMixin:
d2h_event = None
next_batch_result = None
recv_reqs = self.request_receiver.recv_requests()
self.process_input_requests(recv_reqs)
recv_reqs = self.ingest_requests()
if not self.pp_group.is_last_rank:
self._pp_commit_comm_work(self.send_req_work)
@@ -386,8 +384,7 @@ class SchedulerPPMixin:
d2h_event = None
next_batch_result = None
recv_reqs = self.request_receiver.recv_requests()
self.process_input_requests(recv_reqs)
recv_reqs = self.ingest_requests()
if not self.pp_group.is_last_rank:
self._pp_commit_comm_work(self.send_req_work)
@@ -114,8 +114,7 @@ class SchedulerMultiplexMixin:
while True:
with torch.cuda.stream(decode_stream):
set_pdmux_status(False)
recv_reqs = self.request_receiver.recv_requests()
self.process_input_requests(recv_reqs)
self.ingest_requests()
running_batch = self.running_batch
with torch.cuda.stream(prefill_stream):
@@ -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: [],
)