[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):