Add SchedulerRequestReceiver and route request-ingress state through it (#25609)
This commit is contained in:
@@ -1573,7 +1573,9 @@ class SchedulerDisaggregationDecodeMixin:
|
|||||||
|
|
||||||
while True:
|
while True:
|
||||||
# Receive requests
|
# Receive requests
|
||||||
recv_reqs = self.recv_requests()
|
recv_reqs = self.recv_requests(
|
||||||
|
self.request_receiver,
|
||||||
|
)
|
||||||
self.process_input_requests(recv_reqs)
|
self.process_input_requests(recv_reqs)
|
||||||
self.process_decode_queue()
|
self.process_decode_queue()
|
||||||
if self._engine_paused:
|
if self._engine_paused:
|
||||||
@@ -1601,7 +1603,9 @@ class SchedulerDisaggregationDecodeMixin:
|
|||||||
|
|
||||||
while True:
|
while True:
|
||||||
# Receive requests
|
# Receive requests
|
||||||
recv_reqs = self.recv_requests()
|
recv_reqs = self.recv_requests(
|
||||||
|
self.request_receiver,
|
||||||
|
)
|
||||||
self.process_input_requests(recv_reqs)
|
self.process_input_requests(recv_reqs)
|
||||||
self.process_decode_queue()
|
self.process_decode_queue()
|
||||||
if self._engine_paused:
|
if self._engine_paused:
|
||||||
|
|||||||
@@ -395,7 +395,9 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
|
|
||||||
while True:
|
while True:
|
||||||
# Receive requests
|
# Receive requests
|
||||||
recv_reqs = self.recv_requests()
|
recv_reqs = self.recv_requests(
|
||||||
|
self.request_receiver,
|
||||||
|
)
|
||||||
self.process_input_requests(recv_reqs)
|
self.process_input_requests(recv_reqs)
|
||||||
self.waiting_queue.extend(
|
self.waiting_queue.extend(
|
||||||
self.disagg_prefill_bootstrap_queue.pop_bootstrapped()
|
self.disagg_prefill_bootstrap_queue.pop_bootstrapped()
|
||||||
@@ -428,7 +430,9 @@ class SchedulerDisaggregationPrefillMixin:
|
|||||||
|
|
||||||
while True:
|
while True:
|
||||||
# Receive requests
|
# Receive requests
|
||||||
recv_reqs = self.recv_requests()
|
recv_reqs = self.recv_requests(
|
||||||
|
self.request_receiver,
|
||||||
|
)
|
||||||
self.process_input_requests(recv_reqs)
|
self.process_input_requests(recv_reqs)
|
||||||
self.waiting_queue.extend(
|
self.waiting_queue.extend(
|
||||||
self.disagg_prefill_bootstrap_queue.pop_bootstrapped()
|
self.disagg_prefill_bootstrap_queue.pop_bootstrapped()
|
||||||
|
|||||||
@@ -168,7 +168,9 @@ class SchedulerMlxOverlapMixin:
|
|||||||
)
|
)
|
||||||
|
|
||||||
while True:
|
while True:
|
||||||
recv_reqs = self.recv_requests()
|
recv_reqs = self.recv_requests(
|
||||||
|
self.request_receiver,
|
||||||
|
)
|
||||||
self.process_input_requests(recv_reqs)
|
self.process_input_requests(recv_reqs)
|
||||||
if self._engine_paused:
|
if self._engine_paused:
|
||||||
continue
|
continue
|
||||||
|
|||||||
@@ -168,6 +168,9 @@ from sglang.srt.managers.schedule_policy import (
|
|||||||
PrefillAdder,
|
PrefillAdder,
|
||||||
SchedulePolicy,
|
SchedulePolicy,
|
||||||
)
|
)
|
||||||
|
from sglang.srt.managers.scheduler_components.request_receiver import (
|
||||||
|
SchedulerRequestReceiver,
|
||||||
|
)
|
||||||
from sglang.srt.managers.scheduler_dp_attn_mixin import SchedulerDPAttnMixin
|
from sglang.srt.managers.scheduler_dp_attn_mixin import SchedulerDPAttnMixin
|
||||||
from sglang.srt.managers.scheduler_input_blocker import SchedulerInputBlocker
|
from sglang.srt.managers.scheduler_input_blocker import SchedulerInputBlocker
|
||||||
from sglang.srt.managers.scheduler_output_processor_mixin import (
|
from sglang.srt.managers.scheduler_output_processor_mixin import (
|
||||||
@@ -563,6 +566,29 @@ class Scheduler(
|
|||||||
# Init the grammar backend for constrained generation
|
# Init the grammar backend for constrained generation
|
||||||
self.grammar_manager = GrammarManager(self)
|
self.grammar_manager = GrammarManager(self)
|
||||||
|
|
||||||
|
self.request_receiver = SchedulerRequestReceiver(
|
||||||
|
recv_from_tokenizer=self.recv_from_tokenizer,
|
||||||
|
recv_from_rpc=self.recv_from_rpc,
|
||||||
|
recv_skipper=self.recv_skipper,
|
||||||
|
input_blocker=self.input_blocker,
|
||||||
|
mm_receiver=self.mm_receiver,
|
||||||
|
ps=self.ps,
|
||||||
|
tp_group=self.tp_group,
|
||||||
|
tp_cpu_group=self.tp_cpu_group,
|
||||||
|
attn_tp_group=self.attn_tp_group,
|
||||||
|
attn_tp_cpu_group=self.attn_tp_cpu_group,
|
||||||
|
attn_cp_group=self.attn_cp_group,
|
||||||
|
attn_cp_cpu_group=self.attn_cp_cpu_group,
|
||||||
|
world_group=self.world_group,
|
||||||
|
server_args=self.server_args,
|
||||||
|
model_config=self.model_config,
|
||||||
|
max_recv_per_poll=self.max_recv_per_poll,
|
||||||
|
stream_output=self.stream_output,
|
||||||
|
get_last_forward_mode=lambda: (
|
||||||
|
self.last_batch.forward_mode if self.last_batch is not None else None
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
self.is_initializing = False
|
self.is_initializing = False
|
||||||
|
|
||||||
def init_zbal_on_npu(self):
|
def init_zbal_on_npu(self):
|
||||||
@@ -1362,7 +1388,9 @@ class Scheduler(
|
|||||||
"""A normal scheduler loop."""
|
"""A normal scheduler loop."""
|
||||||
while True:
|
while True:
|
||||||
# Receive requests
|
# Receive requests
|
||||||
recv_reqs = self.recv_requests()
|
recv_reqs = self.recv_requests(
|
||||||
|
self.request_receiver,
|
||||||
|
)
|
||||||
self.process_input_requests(recv_reqs)
|
self.process_input_requests(recv_reqs)
|
||||||
if self._engine_paused:
|
if self._engine_paused:
|
||||||
continue
|
continue
|
||||||
@@ -1398,7 +1426,9 @@ class Scheduler(
|
|||||||
|
|
||||||
while True:
|
while True:
|
||||||
# Receive requests
|
# Receive requests
|
||||||
recv_reqs = self.recv_requests()
|
recv_reqs = self.recv_requests(
|
||||||
|
self.request_receiver,
|
||||||
|
)
|
||||||
self.process_input_requests(recv_reqs)
|
self.process_input_requests(recv_reqs)
|
||||||
if self._engine_paused:
|
if self._engine_paused:
|
||||||
continue
|
continue
|
||||||
@@ -1472,21 +1502,22 @@ class Scheduler(
|
|||||||
|
|
||||||
return disable_overlap_for_batch or need_grammar_sync
|
return disable_overlap_for_batch or need_grammar_sync
|
||||||
|
|
||||||
def recv_limit_reached(self, num_recv_reqs: int) -> bool:
|
@staticmethod
|
||||||
|
def recv_limit_reached(
|
||||||
|
self: "SchedulerRequestReceiver", num_recv_reqs: int
|
||||||
|
) -> bool:
|
||||||
if self.max_recv_per_poll < 0:
|
if self.max_recv_per_poll < 0:
|
||||||
return False
|
return False
|
||||||
return num_recv_reqs >= self.max_recv_per_poll
|
return num_recv_reqs >= self.max_recv_per_poll
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
def recv_requests(
|
def recv_requests(
|
||||||
self,
|
self: "SchedulerRequestReceiver",
|
||||||
) -> 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."""
|
||||||
|
|
||||||
if self.recv_skipper is not None:
|
if self.recv_skipper is not None:
|
||||||
last_forward_mode = (
|
if not self.recv_skipper.handle(self.get_last_forward_mode()):
|
||||||
self.last_batch.forward_mode if self.last_batch is not None else None
|
|
||||||
)
|
|
||||||
if not self.recv_skipper.handle(last_forward_mode):
|
|
||||||
return []
|
return []
|
||||||
|
|
||||||
if self.ps.pp_rank == 0:
|
if self.ps.pp_rank == 0:
|
||||||
@@ -1495,7 +1526,7 @@ class Scheduler(
|
|||||||
|
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
if self.recv_limit_reached(len(recv_reqs)):
|
if Scheduler.recv_limit_reached(self, len(recv_reqs)):
|
||||||
break
|
break
|
||||||
recv_req = self.recv_from_tokenizer.recv_pyobj(zmq.NOBLOCK)
|
recv_req = self.recv_from_tokenizer.recv_pyobj(zmq.NOBLOCK)
|
||||||
except zmq.ZMQError:
|
except zmq.ZMQError:
|
||||||
@@ -1504,7 +1535,7 @@ class Scheduler(
|
|||||||
|
|
||||||
while True:
|
while True:
|
||||||
try:
|
try:
|
||||||
if self.recv_limit_reached(len(recv_reqs)):
|
if Scheduler.recv_limit_reached(self, len(recv_reqs)):
|
||||||
break
|
break
|
||||||
recv_rpc = self.recv_from_rpc.recv_pyobj(zmq.NOBLOCK)
|
recv_rpc = self.recv_from_rpc.recv_pyobj(zmq.NOBLOCK)
|
||||||
except zmq.ZMQError:
|
except zmq.ZMQError:
|
||||||
@@ -1530,7 +1561,9 @@ class Scheduler(
|
|||||||
|
|
||||||
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 = Scheduler._split_work_and_control_reqs(
|
||||||
|
self, recv_reqs
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
work_reqs = None
|
work_reqs = None
|
||||||
control_reqs = None
|
control_reqs = None
|
||||||
@@ -1634,7 +1667,8 @@ class Scheduler(
|
|||||||
|
|
||||||
return recv_reqs
|
return recv_reqs
|
||||||
|
|
||||||
def _split_work_and_control_reqs(self, recv_reqs: List):
|
@staticmethod
|
||||||
|
def _split_work_and_control_reqs(self: "SchedulerRequestReceiver", recv_reqs: List):
|
||||||
work_reqs = [
|
work_reqs = [
|
||||||
req
|
req
|
||||||
for req in recv_reqs
|
for req in recv_reqs
|
||||||
|
|||||||
@@ -0,0 +1,33 @@
|
|||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from typing import TYPE_CHECKING, Any, Callable, Optional
|
||||||
|
|
||||||
|
import zmq
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
from sglang.srt.configs.model_config import ModelConfig
|
||||||
|
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
|
||||||
|
from sglang.srt.server_args import ServerArgs
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass(kw_only=True, slots=True, frozen=True)
|
||||||
|
class SchedulerRequestReceiver:
|
||||||
|
recv_from_tokenizer: zmq.Socket
|
||||||
|
recv_from_rpc: Optional[zmq.Socket]
|
||||||
|
recv_skipper: Any
|
||||||
|
input_blocker: Any
|
||||||
|
mm_receiver: Any
|
||||||
|
ps: "ParallelState"
|
||||||
|
tp_group: Any
|
||||||
|
tp_cpu_group: Any
|
||||||
|
attn_tp_group: Any
|
||||||
|
attn_tp_cpu_group: Any
|
||||||
|
attn_cp_group: Any
|
||||||
|
attn_cp_cpu_group: Any
|
||||||
|
world_group: Any
|
||||||
|
server_args: "ServerArgs"
|
||||||
|
model_config: "ModelConfig"
|
||||||
|
max_recv_per_poll: int
|
||||||
|
stream_output: Callable[..., None]
|
||||||
|
get_last_forward_mode: Callable[[], Any]
|
||||||
@@ -80,7 +80,9 @@ 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.recv_requests()
|
recv_reqs = self.recv_requests(
|
||||||
|
self.request_receiver,
|
||||||
|
)
|
||||||
self.process_input_requests(recv_reqs)
|
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)
|
||||||
@@ -214,7 +216,9 @@ class SchedulerPPMixin:
|
|||||||
d2h_event = None
|
d2h_event = None
|
||||||
next_batch_result = None
|
next_batch_result = None
|
||||||
|
|
||||||
recv_reqs = self.recv_requests()
|
recv_reqs = self.recv_requests(
|
||||||
|
self.request_receiver,
|
||||||
|
)
|
||||||
self.process_input_requests(recv_reqs)
|
self.process_input_requests(recv_reqs)
|
||||||
|
|
||||||
if not self.pp_group.is_last_rank:
|
if not self.pp_group.is_last_rank:
|
||||||
@@ -360,7 +364,9 @@ class SchedulerPPMixin:
|
|||||||
d2h_event = None
|
d2h_event = None
|
||||||
next_batch_result = None
|
next_batch_result = None
|
||||||
|
|
||||||
recv_reqs = self.recv_requests()
|
recv_reqs = self.recv_requests(
|
||||||
|
self.request_receiver,
|
||||||
|
)
|
||||||
self.process_input_requests(recv_reqs)
|
self.process_input_requests(recv_reqs)
|
||||||
|
|
||||||
if not self.pp_group.is_last_rank:
|
if not self.pp_group.is_last_rank:
|
||||||
|
|||||||
@@ -110,7 +110,9 @@ 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.recv_requests()
|
recv_reqs = self.recv_requests(
|
||||||
|
self.request_receiver,
|
||||||
|
)
|
||||||
self.process_input_requests(recv_reqs)
|
self.process_input_requests(recv_reqs)
|
||||||
|
|
||||||
with torch.cuda.stream(prefill_stream):
|
with torch.cuda.stream(prefill_stream):
|
||||||
|
|||||||
Reference in New Issue
Block a user