[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:
|
if not self._engine_paused:
|
||||||
self.disagg_decode_prealloc_queue.prefetch_prefill_dp_rank_queries()
|
self.disagg_decode_prealloc_queue.prefetch_prefill_dp_rank_queries()
|
||||||
# Receive requests
|
# Receive requests
|
||||||
recv_reqs = self.request_receiver.recv_requests()
|
self.ingest_requests()
|
||||||
self.process_input_requests(recv_reqs)
|
|
||||||
if self._engine_paused:
|
if self._engine_paused:
|
||||||
self._record_scheduler_state_for_paused_engine()
|
self._record_scheduler_state_for_paused_engine()
|
||||||
continue
|
continue
|
||||||
@@ -2548,8 +2547,7 @@ class SchedulerDisaggregationDecodeMixin:
|
|||||||
if not self._engine_paused:
|
if not self._engine_paused:
|
||||||
self.disagg_decode_prealloc_queue.prefetch_prefill_dp_rank_queries()
|
self.disagg_decode_prealloc_queue.prefetch_prefill_dp_rank_queries()
|
||||||
# Receive requests
|
# Receive requests
|
||||||
recv_reqs = self.request_receiver.recv_requests()
|
self.ingest_requests()
|
||||||
self.process_input_requests(recv_reqs)
|
|
||||||
if self._engine_paused:
|
if self._engine_paused:
|
||||||
self._record_scheduler_state_for_paused_engine()
|
self._record_scheduler_state_for_paused_engine()
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -599,8 +599,7 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
"""A normal scheduler loop for prefill worker in disaggregation mode."""
|
"""A normal scheduler loop for prefill worker in disaggregation mode."""
|
||||||
while True:
|
while True:
|
||||||
# Receive requests
|
# Receive requests
|
||||||
recv_reqs = self.request_receiver.recv_requests()
|
self.ingest_requests()
|
||||||
self.process_input_requests(recv_reqs)
|
|
||||||
if self._engine_paused:
|
if self._engine_paused:
|
||||||
self._record_scheduler_state_for_paused_engine()
|
self._record_scheduler_state_for_paused_engine()
|
||||||
continue
|
continue
|
||||||
@@ -639,8 +638,7 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
|
|
||||||
while True:
|
while True:
|
||||||
# Receive requests
|
# Receive requests
|
||||||
recv_reqs = self.request_receiver.recv_requests()
|
self.ingest_requests()
|
||||||
self.process_input_requests(recv_reqs)
|
|
||||||
if self._engine_paused:
|
if self._engine_paused:
|
||||||
self._record_scheduler_state_for_paused_engine()
|
self._record_scheduler_state_for_paused_engine()
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -210,8 +210,7 @@ class SchedulerMlxOverlapMixin:
|
|||||||
mx.synchronize()
|
mx.synchronize()
|
||||||
break
|
break
|
||||||
|
|
||||||
recv_reqs = self.request_receiver.recv_requests()
|
self.ingest_requests()
|
||||||
self.process_input_requests(recv_reqs)
|
|
||||||
if self._engine_paused:
|
if self._engine_paused:
|
||||||
self._record_scheduler_state_for_paused_engine()
|
self._record_scheduler_state_for_paused_engine()
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -1895,10 +1895,7 @@ class Scheduler(
|
|||||||
break
|
break
|
||||||
|
|
||||||
# Receive requests
|
# Receive requests
|
||||||
recv_reqs = self.request_receiver.recv_requests()
|
self.ingest_requests()
|
||||||
if recv_reqs:
|
|
||||||
self.metrics_reporter.record_scheduler_active()
|
|
||||||
self.process_input_requests(recv_reqs)
|
|
||||||
if self._engine_paused:
|
if self._engine_paused:
|
||||||
self._record_scheduler_state_for_paused_engine()
|
self._record_scheduler_state_for_paused_engine()
|
||||||
continue
|
continue
|
||||||
@@ -1942,10 +1939,7 @@ class Scheduler(
|
|||||||
break
|
break
|
||||||
|
|
||||||
# Receive requests
|
# Receive requests
|
||||||
recv_reqs = self.request_receiver.recv_requests()
|
self.ingest_requests()
|
||||||
if recv_reqs:
|
|
||||||
self.metrics_reporter.record_scheduler_active()
|
|
||||||
self.process_input_requests(recv_reqs)
|
|
||||||
if self._engine_paused:
|
if self._engine_paused:
|
||||||
self._record_scheduler_state_for_paused_engine()
|
self._record_scheduler_state_for_paused_engine()
|
||||||
continue
|
continue
|
||||||
@@ -2051,6 +2045,25 @@ class Scheduler(
|
|||||||
for prev_batch, prev_result in self.result_queue:
|
for prev_batch, prev_result in self.result_queue:
|
||||||
self.batch_result_processor.advance_grammar_fsm(prev_result, prev_batch)
|
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)
|
@scheduler_stage_method(SCHEDULER_STAGE_PROCESS_REQUESTS)
|
||||||
def process_input_requests(self, recv_reqs: List):
|
def process_input_requests(self, recv_reqs: List):
|
||||||
now = time.monotonic()
|
now = time.monotonic()
|
||||||
@@ -2278,7 +2291,6 @@ class Scheduler(
|
|||||||
get_last_batch=lambda: self.last_batch,
|
get_last_batch=lambda: self.last_batch,
|
||||||
scripted_scheduler_hook=self.scripted_scheduler_hook,
|
scripted_scheduler_hook=self.scripted_scheduler_hook,
|
||||||
scheduler_stage_metrics=self.scheduler_stage_metrics,
|
scheduler_stage_metrics=self.scheduler_stage_metrics,
|
||||||
poll_timeout_aborts=self._poll_timeout_aborts,
|
|
||||||
)
|
)
|
||||||
|
|
||||||
def init_dp_attn_adapter(self) -> None:
|
def init_dp_attn_adapter(self) -> None:
|
||||||
|
|||||||
@@ -79,9 +79,6 @@ class SchedulerRequestReceiver:
|
|||||||
get_last_batch: Callable[[], Any]
|
get_last_batch: Callable[[], Any]
|
||||||
scripted_scheduler_hook: Optional[ScriptedSchedulerHook] = None
|
scripted_scheduler_hook: Optional[ScriptedSchedulerHook] = None
|
||||||
scheduler_stage_metrics: Optional[SchedulerStageMetricsRecorder] = 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:
|
def recv_limit_reached(self, num_recv_reqs: int) -> bool:
|
||||||
if self.max_recv_per_poll < 0:
|
if self.max_recv_per_poll < 0:
|
||||||
@@ -90,9 +87,13 @@ class SchedulerRequestReceiver:
|
|||||||
|
|
||||||
@scheduler_stage_method(SCHEDULER_STAGE_RECV_REQUESTS)
|
@scheduler_stage_method(SCHEDULER_STAGE_RECV_REQUESTS)
|
||||||
def recv_requests(
|
def recv_requests(
|
||||||
self,
|
self, local_reqs: Optional[List[AbortReq]] = None
|
||||||
) -> List[Union[TokenizedGenerateReqInput, TokenizedEmbeddingReqInput, Any]]:
|
) -> 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:
|
if self.scripted_scheduler_hook is not None:
|
||||||
self.scripted_scheduler_hook.step()
|
self.scripted_scheduler_hook.step()
|
||||||
@@ -106,16 +107,6 @@ class SchedulerRequestReceiver:
|
|||||||
if self.input_blocker is not None:
|
if self.input_blocker is not None:
|
||||||
recv_reqs = self.input_blocker.handle(recv_reqs)
|
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)
|
recv_reqs = self._broadcast_reqs_across_ranks(recv_reqs, local_reqs)
|
||||||
|
|
||||||
if self.ps.pp_rank == 0:
|
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_first_rank_mb_id = (mb_id + self.ps.pp_size) % self.pp_loop_size
|
||||||
next_mb_id = (mb_id + 1) % self.pp_loop_size
|
next_mb_id = (mb_id + 1) % self.pp_loop_size
|
||||||
with torch.profiler.record_function("recv_requests"):
|
with torch.profiler.record_function("recv_requests"):
|
||||||
recv_reqs = self.request_receiver.recv_requests()
|
recv_reqs = self.ingest_requests()
|
||||||
self.process_input_requests(recv_reqs)
|
|
||||||
if not self.pp_group.is_last_rank:
|
if not self.pp_group.is_last_rank:
|
||||||
self._pp_commit_comm_work(self.send_req_work)
|
self._pp_commit_comm_work(self.send_req_work)
|
||||||
with torch.profiler.record_function("send_reqs_to_next_stage"):
|
with torch.profiler.record_function("send_reqs_to_next_stage"):
|
||||||
@@ -232,8 +231,7 @@ class SchedulerPPMixin:
|
|||||||
d2h_event = None
|
d2h_event = None
|
||||||
next_batch_result = None
|
next_batch_result = None
|
||||||
|
|
||||||
recv_reqs = self.request_receiver.recv_requests()
|
recv_reqs = self.ingest_requests()
|
||||||
self.process_input_requests(recv_reqs)
|
|
||||||
|
|
||||||
if not self.pp_group.is_last_rank:
|
if not self.pp_group.is_last_rank:
|
||||||
self._pp_commit_comm_work(self.send_req_work)
|
self._pp_commit_comm_work(self.send_req_work)
|
||||||
@@ -386,8 +384,7 @@ class SchedulerPPMixin:
|
|||||||
d2h_event = None
|
d2h_event = None
|
||||||
next_batch_result = None
|
next_batch_result = None
|
||||||
|
|
||||||
recv_reqs = self.request_receiver.recv_requests()
|
recv_reqs = self.ingest_requests()
|
||||||
self.process_input_requests(recv_reqs)
|
|
||||||
|
|
||||||
if not self.pp_group.is_last_rank:
|
if not self.pp_group.is_last_rank:
|
||||||
self._pp_commit_comm_work(self.send_req_work)
|
self._pp_commit_comm_work(self.send_req_work)
|
||||||
|
|||||||
@@ -114,8 +114,7 @@ class SchedulerMultiplexMixin:
|
|||||||
while True:
|
while True:
|
||||||
with torch.cuda.stream(decode_stream):
|
with torch.cuda.stream(decode_stream):
|
||||||
set_pdmux_status(False)
|
set_pdmux_status(False)
|
||||||
recv_reqs = self.request_receiver.recv_requests()
|
self.ingest_requests()
|
||||||
self.process_input_requests(recv_reqs)
|
|
||||||
running_batch = self.running_batch
|
running_batch = self.running_batch
|
||||||
|
|
||||||
with torch.cuda.stream(prefill_stream):
|
with torch.cuda.stream(prefill_stream):
|
||||||
|
|||||||
@@ -1154,8 +1154,7 @@ class TestMlxOverlapScheduler(unittest.TestCase):
|
|||||||
raise _StopLoop
|
raise _StopLoop
|
||||||
|
|
||||||
scheduler = SchedulerMlxOverlapMixin.__new__(SchedulerMlxOverlapMixin)
|
scheduler = SchedulerMlxOverlapMixin.__new__(SchedulerMlxOverlapMixin)
|
||||||
scheduler.request_receiver = SimpleNamespace(recv_requests=lambda: [])
|
scheduler.ingest_requests = lambda: []
|
||||||
scheduler.process_input_requests = lambda recv_reqs: None
|
|
||||||
scheduler.gracefully_exit = False
|
scheduler.gracefully_exit = False
|
||||||
scheduler._engine_paused = False
|
scheduler._engine_paused = False
|
||||||
scheduler.forward_ct = 0
|
scheduler.forward_ct = 0
|
||||||
|
|||||||
@@ -118,7 +118,7 @@ class TestOverlapLoopStampsLaunchTs(unittest.TestCase):
|
|||||||
scheduler._engine_paused = False
|
scheduler._engine_paused = False
|
||||||
scheduler.waiting_queue = []
|
scheduler.waiting_queue = []
|
||||||
scheduler.result_queue = deque()
|
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 = MagicMock()
|
||||||
result.next_token_ids = None
|
result.next_token_ids = None
|
||||||
scheduler.tp_worker.finalize_mlx_result.return_value = result
|
scheduler.tp_worker.finalize_mlx_result.return_value = result
|
||||||
@@ -269,7 +269,7 @@ class TestOverlapLoopGracefulExit(unittest.TestCase):
|
|||||||
scheduler._engine_paused = False
|
scheduler._engine_paused = False
|
||||||
scheduler.waiting_queue = []
|
scheduler.waiting_queue = []
|
||||||
scheduler.result_queue = deque()
|
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
|
# Model handle_shutdown: processing a non-empty recv batch (the
|
||||||
# ShutdownReq) flips the flag; the loop must notice at the top of the
|
# ShutdownReq) flips the flag; the loop must notice at the top of the
|
||||||
# next iteration instead of polling forever.
|
# next iteration instead of polling forever.
|
||||||
@@ -296,7 +296,7 @@ class TestOverlapLoopGracefulExit(unittest.TestCase):
|
|||||||
) as synchronize:
|
) as synchronize:
|
||||||
SchedulerMlxOverlapMixin.event_loop_overlap_mlx(scheduler)
|
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()
|
synchronize.assert_called_once_with()
|
||||||
|
|
||||||
def test_loop_exits_when_shutdown_arrives_while_paused(self):
|
def test_loop_exits_when_shutdown_arrives_while_paused(self):
|
||||||
@@ -317,7 +317,7 @@ class TestOverlapLoopGracefulExit(unittest.TestCase):
|
|||||||
) as synchronize:
|
) as synchronize:
|
||||||
SchedulerMlxOverlapMixin.event_loop_overlap_mlx(scheduler)
|
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()
|
scheduler.get_next_batch_to_run.assert_not_called()
|
||||||
synchronize.assert_called_once_with()
|
synchronize.assert_called_once_with()
|
||||||
|
|
||||||
|
|||||||
@@ -96,12 +96,9 @@ class TestMambaBoundaryMaskReuse(unittest.TestCase):
|
|||||||
|
|
||||||
scheduler = Scheduler.__new__(Scheduler)
|
scheduler = Scheduler.__new__(Scheduler)
|
||||||
scheduler.gracefully_exit = False
|
scheduler.gracefully_exit = False
|
||||||
scheduler.request_receiver = MagicMock()
|
scheduler.ingest_requests = MagicMock(
|
||||||
scheduler.request_receiver.recv_requests.side_effect = [
|
side_effect=[[], [], StopIteration]
|
||||||
[],
|
)
|
||||||
[],
|
|
||||||
StopIteration,
|
|
||||||
]
|
|
||||||
scheduler.process_input_requests = MagicMock()
|
scheduler.process_input_requests = MagicMock()
|
||||||
scheduler._engine_paused = False
|
scheduler._engine_paused = False
|
||||||
scheduler.running_batch = batch
|
scheduler.running_batch = batch
|
||||||
|
|||||||
@@ -121,7 +121,6 @@ def _receiver(tp_size: int = 1) -> SchedulerRequestReceiver:
|
|||||||
max_recv_per_poll=-1,
|
max_recv_per_poll=-1,
|
||||||
stream_output=lambda *args, **kwargs: None,
|
stream_output=lambda *args, **kwargs: None,
|
||||||
get_last_batch=lambda: 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,
|
max_recv_per_poll=-1,
|
||||||
stream_output=lambda *args, **kwargs: None,
|
stream_output=lambda *args, **kwargs: None,
|
||||||
get_last_batch=lambda: None,
|
get_last_batch=lambda: None,
|
||||||
poll_timeout_aborts=lambda: [],
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user