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()): if not self.recv_skipper.handle(self.get_last_forward_mode()):
return [] 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.pp_rank == 0:
if self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0: if self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0:
recv_reqs = [] recv_reqs = []
@@ -106,10 +120,9 @@ class SchedulerRequestReceiver:
) )
else: else:
recv_reqs = None recv_reqs = None
return recv_reqs
if self.input_blocker is not None: def _broadcast_reqs_across_ranks(self, recv_reqs: Optional[List]) -> List:
recv_reqs = self.input_blocker.handle(recv_reqs)
if self.server_args.enable_dp_attention: if self.server_args.enable_dp_attention:
if self.ps.attn_tp_rank == 0 and self.ps.attn_cp_rank == 0: 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) work_reqs, control_reqs = self._split_work_and_control_reqs(recv_reqs)
@@ -169,7 +182,9 @@ class SchedulerRequestReceiver:
self.tp_cpu_group, self.tp_cpu_group,
src=self.tp_group.ranks[0], 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 # Process MM requests under EPD-disaggregation mode
if ( if (
self.ps.pp_rank == 0 self.ps.pp_rank == 0
@@ -185,7 +200,9 @@ class SchedulerRequestReceiver:
) )
prepare_abort(req, error_msg, status_code=status_code) prepare_abort(req, error_msg, status_code=status_code)
self.stream_output([req], req.return_logprob) 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, # Unwrap shared memory features AFTER all broadcasts complete,
# so that ShmPointerMMData metadata (not full tensor data) is what # so that ShmPointerMMData metadata (not full tensor data) is what
# gets serialized during broadcast_pyobj. # gets serialized during broadcast_pyobj.
@@ -214,8 +231,6 @@ class SchedulerRequestReceiver:
for req in recv_reqs: for req in recv_reqs:
unwrap_shm_features(req) unwrap_shm_features(req)
return recv_reqs
def _split_work_and_control_reqs(self, recv_reqs: List): def _split_work_and_control_reqs(self, recv_reqs: List):
work_reqs = [ work_reqs = [
req req