[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):
|
||||
|
||||
@@ -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: [],
|
||||
)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user