diff --git a/python/sglang/srt/disaggregation/decode.py b/python/sglang/srt/disaggregation/decode.py index 11c1494c1..41a1a1ec2 100644 --- a/python/sglang/srt/disaggregation/decode.py +++ b/python/sglang/srt/disaggregation/decode.py @@ -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 diff --git a/python/sglang/srt/disaggregation/prefill.py b/python/sglang/srt/disaggregation/prefill.py index afe27cc58..20e616590 100644 --- a/python/sglang/srt/disaggregation/prefill.py +++ b/python/sglang/srt/disaggregation/prefill.py @@ -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 diff --git a/python/sglang/srt/hardware_backend/mlx/scheduler_mixin.py b/python/sglang/srt/hardware_backend/mlx/scheduler_mixin.py index 850b2da9f..62d602523 100644 --- a/python/sglang/srt/hardware_backend/mlx/scheduler_mixin.py +++ b/python/sglang/srt/hardware_backend/mlx/scheduler_mixin.py @@ -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 diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index f408205aa..44c7deb00 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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: diff --git a/python/sglang/srt/managers/scheduler_components/request_receiver.py b/python/sglang/srt/managers/scheduler_components/request_receiver.py index 49d2ba7a9..5970fb768 100644 --- a/python/sglang/srt/managers/scheduler_components/request_receiver.py +++ b/python/sglang/srt/managers/scheduler_components/request_receiver.py @@ -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: diff --git a/python/sglang/srt/managers/scheduler_pp_mixin.py b/python/sglang/srt/managers/scheduler_pp_mixin.py index de99cd258..c852d2e83 100644 --- a/python/sglang/srt/managers/scheduler_pp_mixin.py +++ b/python/sglang/srt/managers/scheduler_pp_mixin.py @@ -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) diff --git a/python/sglang/srt/multiplex/multiplexing_mixin.py b/python/sglang/srt/multiplex/multiplexing_mixin.py index 253d42b8f..a52a736bb 100644 --- a/python/sglang/srt/multiplex/multiplexing_mixin.py +++ b/python/sglang/srt/multiplex/multiplexing_mixin.py @@ -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): diff --git a/test/registered/unit/hardware_backend/mlx/test_attention_patching.py b/test/registered/unit/hardware_backend/mlx/test_attention_patching.py index 6ecda2269..47b64f848 100644 --- a/test/registered/unit/hardware_backend/mlx/test_attention_patching.py +++ b/test/registered/unit/hardware_backend/mlx/test_attention_patching.py @@ -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 diff --git a/test/registered/unit/hardware_backend/mlx/test_scheduler_mixin.py b/test/registered/unit/hardware_backend/mlx/test_scheduler_mixin.py index 6e4b9cbbd..f94872db1 100644 --- a/test/registered/unit/hardware_backend/mlx/test_scheduler_mixin.py +++ b/test/registered/unit/hardware_backend/mlx/test_scheduler_mixin.py @@ -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() diff --git a/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py b/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py index d9e90efe4..53fa3b497 100644 --- a/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py +++ b/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py @@ -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 diff --git a/test/registered/unit/managers/test_mm_shm_error_consensus.py b/test/registered/unit/managers/test_mm_shm_error_consensus.py index 110071037..51d951709 100644 --- a/test/registered/unit/managers/test_mm_shm_error_consensus.py +++ b/test/registered/unit/managers/test_mm_shm_error_consensus.py @@ -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: [], ) diff --git a/test/registered/unit/managers/test_pp_cp_rank_offsets.py b/test/registered/unit/managers/test_pp_cp_rank_offsets.py index 21e61d723..bea857174 100644 --- a/test/registered/unit/managers/test_pp_cp_rank_offsets.py +++ b/test/registered/unit/managers/test_pp_cp_rank_offsets.py @@ -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: [], )