Split the request-reception loop into smaller phases (#25703)

This commit is contained in:
fzyzcjy
2026-05-19 09:12:47 +08:00
committed by GitHub
parent d90bc65e30
commit e8e55bb19b
@@ -71,6 +71,20 @@ class SchedulerRequestReceiver:
if not self.recv_skipper.handle(self.get_last_forward_mode()):
return []
recv_reqs = self._pull_raw_reqs()
if self.input_blocker is not None:
recv_reqs = self.input_blocker.handle(recv_reqs)
recv_reqs = self._broadcast_reqs_across_ranks(recv_reqs)
recv_reqs = self._apply_mm_receiver(recv_reqs)
self._finalize_shm_features(recv_reqs)
return recv_reqs
def _pull_raw_reqs(self) -> Optional[List]:
if self.ps.pp_rank == 0:
if self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0:
recv_reqs = []
@@ -106,10 +120,9 @@ class SchedulerRequestReceiver:
)
else:
recv_reqs = None
return recv_reqs
if self.input_blocker is not None:
recv_reqs = self.input_blocker.handle(recv_reqs)
def _broadcast_reqs_across_ranks(self, recv_reqs: Optional[List]) -> List:
if self.server_args.enable_dp_attention:
if self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0:
work_reqs, control_reqs = self._split_work_and_control_reqs(recv_reqs)
@@ -169,7 +182,9 @@ class SchedulerRequestReceiver:
self.tp_cpu_group,
src=self.tp_group.ranks[0],
)
return recv_reqs
def _apply_mm_receiver(self, recv_reqs: List) -> List:
# Process MM requests under EPD-disaggregation mode
if (
self.ps.pp_rank == 0
@@ -185,7 +200,9 @@ class SchedulerRequestReceiver:
)
prepare_abort(req, error_msg, status_code=status_code)
self.stream_output([req], req.return_logprob)
return recv_reqs
def _finalize_shm_features(self, recv_reqs: Optional[List]) -> None:
# Unwrap shared memory features AFTER all broadcasts complete,
# so that ShmPointerMMData metadata (not full tensor data) is what
# gets serialized during broadcast_pyobj.
@@ -214,8 +231,6 @@ class SchedulerRequestReceiver:
for req in recv_reqs:
unwrap_shm_features(req)
return recv_reqs
def _split_work_and_control_reqs(self, recv_reqs: List):
work_reqs = [
req