[Scheduler] Unify per-iteration request intake into ingest_requests() (#38389)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user