diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index c68a235e1..a92cfdc12 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -1896,6 +1896,9 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): can_run_dp_cuda_graph: bool = False can_run_dp_breakable_cuda_graph: bool = False tbo_split_seq_index: Optional[int] = None + # Rank-consistent forward mode for the recv skipper, derived from the MLP + # sync all-gather (the TBO-only `global_forward_mode` is None without TBO). + recv_skipper_forward_mode: Optional[ForwardMode] = None spec_verify_tier_num_tokens: int = -1 # For processing logprobs diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index 4e2bee73b..f04458c2c 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -1726,9 +1726,7 @@ class Scheduler( model_config=self.model_config, max_recv_per_poll=self.max_recv_per_poll, stream_output=lambda *a, **kw: self.output_streamer.stream_output(*a, **kw), - get_last_forward_mode=lambda: ( - self.last_batch.forward_mode if self.last_batch is not None else None - ), + get_last_batch=lambda: self.last_batch, scripted_scheduler_hook=self.scripted_scheduler_hook, ) diff --git a/python/sglang/srt/managers/scheduler_components/dp_attn.py b/python/sglang/srt/managers/scheduler_components/dp_attn.py index a4efe8218..10f18df75 100644 --- a/python/sglang/srt/managers/scheduler_components/dp_attn.py +++ b/python/sglang/srt/managers/scheduler_components/dp_attn.py @@ -11,6 +11,7 @@ from sglang.srt.distributed.parallel_state import get_tp_group from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.environ import envs from sglang.srt.managers.schedule_batch import ScheduleBatch +from sglang.srt.managers.scheduler_recv_skipper import SchedulerRecvSkipper from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator from sglang.srt.mem_cache.base_prefix_cache import BasePrefixCache from sglang.srt.mem_cache.memory_pool import ReqToTokenPool @@ -257,6 +258,15 @@ def prepare_mlp_sync_batch_raw( batch_to_gather, mlp_sync_info, require_mlp_tp_gather, skip_all_gather ) + # Set on `local_batch`, not `batch_to_gather`: for PREBUILT batches the + # scheduler's `last_batch` is the prebuilt batch, not its inner idle batch. + if local_batch is not None and not skip_all_gather: + local_batch.recv_skipper_forward_mode = ( + SchedulerRecvSkipper.derive_forward_mode( + mlp_sync_info.tp0_info[:, 5].tolist() + ) + ) + if _ENABLE_METRICS_DP_ATTENTION and local_batch is not None: local_batch.dp_cooperation_info = mlp_sync_info.dp_cooperation_info diff --git a/python/sglang/srt/managers/scheduler_components/request_receiver.py b/python/sglang/srt/managers/scheduler_components/request_receiver.py index 27a32de18..a27664adc 100644 --- a/python/sglang/srt/managers/scheduler_components/request_receiver.py +++ b/python/sglang/srt/managers/scheduler_components/request_receiver.py @@ -61,7 +61,7 @@ class SchedulerRequestReceiver: model_config: ModelConfig max_recv_per_poll: int stream_output: Callable[..., None] - get_last_forward_mode: Callable[[], Any] + get_last_batch: Callable[[], Any] scripted_scheduler_hook: Optional[ScriptedSchedulerHook] = None def recv_limit_reached(self, num_recv_reqs: int) -> bool: @@ -79,7 +79,7 @@ class SchedulerRequestReceiver: self.scripted_scheduler_hook.step() if self.recv_skipper is not None: - if not self.recv_skipper.handle(self.get_last_forward_mode()): + if not self.recv_skipper.handle(self.get_last_batch()): return [] recv_reqs = self._pull_raw_reqs() diff --git a/python/sglang/srt/managers/scheduler_recv_skipper.py b/python/sglang/srt/managers/scheduler_recv_skipper.py index 69c3e19a5..946e50247 100644 --- a/python/sglang/srt/managers/scheduler_recv_skipper.py +++ b/python/sglang/srt/managers/scheduler_recv_skipper.py @@ -1,7 +1,14 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING, List, Optional + from sglang.srt.environ import envs from sglang.srt.model_executor.forward_batch_info import ForwardMode from sglang.srt.server_args import ServerArgs +if TYPE_CHECKING: + from sglang.srt.managers.schedule_batch import ScheduleBatch + class SchedulerRecvSkipper: @staticmethod @@ -10,9 +17,24 @@ class SchedulerRecvSkipper: return None return SchedulerRecvSkipper(server_args) + @staticmethod + def derive_forward_mode(gathered_modes: List[int]) -> Optional[ForwardMode]: + """Collapse the gathered per-DP-rank forward modes into one weight-table + bucket; the input is rank-identical, so the recv decision is too.""" + active = set(gathered_modes) - { + ForwardMode.IDLE.value, + ForwardMode.PREBUILT.value, + } + if not active: + return None # globally idle: same bucket as "no last batch" + if active - {ForwardMode.DECODE.value, ForwardMode.TARGET_VERIFY.value}: + return ForwardMode.EXTEND # any extend-like rank: prompt recv + if active == {ForwardMode.TARGET_VERIFY.value}: + return ForwardMode.TARGET_VERIFY + return ForwardMode.DECODE + def __init__(self, server_args: ServerArgs): - # Can be supported if needed, but may need e.g. `global_forward_mode` - assert not server_args.enable_dp_attention + self._use_synced_mode = server_args.enable_dp_attention self._counter = 0 self._threshold = server_args.scheduler_recv_interval # All can be tuned if needed @@ -23,11 +45,21 @@ class SchedulerRecvSkipper: None: envs.SGLANG_SCHEDULER_RECV_SKIPPER_WEIGHT_NONE.get(), } - def handle(self, last_forward_mode: ForwardMode): + def _pick_mode(self, last_batch: Optional[ScheduleBatch]) -> Optional[ForwardMode]: + # The recv decision must be identical on every rank in the request + # broadcast. Local modes differ across DP ranks (IDLE vs DECODE), so + # use the rank-consistent mode derived from the MLP sync all-gather. + if last_batch is None: + return None + if self._use_synced_mode: + return last_batch.recv_skipper_forward_mode + return last_batch.forward_mode + + def handle(self, last_batch: Optional[ScheduleBatch]) -> bool: should_recv = False last_weight = self._weight_of_forward_mode.get( - last_forward_mode, self._default_weight + self._pick_mode(last_batch), self._default_weight ) self._counter += last_weight 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 619c13f92..57f03c9a9 100644 --- a/test/registered/unit/managers/test_pp_cp_rank_offsets.py +++ b/test/registered/unit/managers/test_pp_cp_rank_offsets.py @@ -59,7 +59,7 @@ def _make_receiver(ps: ParallelState) -> SchedulerRequestReceiver: model_config=SimpleNamespace(is_multimodal=False), max_recv_per_poll=-1, stream_output=lambda *args, **kwargs: None, - get_last_forward_mode=lambda: None, + get_last_batch=lambda: None, ) diff --git a/test/registered/unit/managers/test_scheduler_recv_skipper.py b/test/registered/unit/managers/test_scheduler_recv_skipper.py new file mode 100644 index 000000000..df3e14e3d --- /dev/null +++ b/test/registered/unit/managers/test_scheduler_recv_skipper.py @@ -0,0 +1,99 @@ +import unittest +from types import SimpleNamespace + +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import CustomTestCase, maybe_stub_sgl_kernel + +maybe_stub_sgl_kernel() + +from sglang.srt.managers.scheduler_recv_skipper import ( # noqa: E402 + SchedulerRecvSkipper, +) +from sglang.srt.model_executor.forward_batch_info import ForwardMode # noqa: E402 + +register_cpu_ci(est_time=2, suite="base-a-test-cpu") + + +def _server_args(interval, enable_dp_attention=False): + return SimpleNamespace( + scheduler_recv_interval=interval, + enable_dp_attention=enable_dp_attention, + ) + + +def _batch(forward_mode, recv_skipper_forward_mode=None): + return SimpleNamespace( + forward_mode=forward_mode, + recv_skipper_forward_mode=recv_skipper_forward_mode, + ) + + +class TestSchedulerRecvSkipper(CustomTestCase): + def test_disabled_at_default_interval(self): + # interval <= 1 disables the skipper entirely. + self.assertIsNone(SchedulerRecvSkipper.maybe_create(_server_args(1))) + + def test_enabled_under_dp_attention(self): + # Regression: the constructor used to assert `not enable_dp_attention`. + skipper = SchedulerRecvSkipper.maybe_create( + _server_args(50, enable_dp_attention=True) + ) + self.assertIsNotNone(skipper) + + def test_no_last_batch_accumulates_slowly(self): + skipper = SchedulerRecvSkipper.maybe_create(_server_args(50)) + self.assertFalse(skipper.handle(None)) + + def test_decode_accumulates_until_threshold(self): + # DECODE weight = 1: recv only every `interval` decode steps. + skipper = SchedulerRecvSkipper.maybe_create(_server_args(3)) + decode = _batch(ForwardMode.DECODE) + self.assertFalse(skipper.handle(decode)) # counter 1 + self.assertFalse(skipper.handle(decode)) # counter 2 + self.assertTrue(skipper.handle(decode)) # counter 3 -> recv, reset + self.assertFalse(skipper.handle(decode)) # counter 1 again + + def test_prefill_triggers_recv_immediately(self): + # Non-decode passes use the large default weight: recv right away. + skipper = SchedulerRecvSkipper.maybe_create(_server_args(50)) + self.assertTrue(skipper.handle(_batch(ForwardMode.EXTEND))) + + def test_dp_uses_synced_mode_not_local(self): + # Local EXTEND (weight 1000) must be ignored in favor of the synced + # DECODE (weight 1); a recv here would mean the local mode leaked in. + skipper = SchedulerRecvSkipper.maybe_create( + _server_args(50, enable_dp_attention=True) + ) + self.assertFalse(skipper.handle(_batch(ForwardMode.EXTEND, ForwardMode.DECODE))) + + def test_dp_synced_extend_triggers_recv(self): + skipper = SchedulerRecvSkipper.maybe_create( + _server_args(50, enable_dp_attention=True) + ) + self.assertTrue(skipper.handle(_batch(ForwardMode.IDLE, ForwardMode.EXTEND))) + + def test_derive_forward_mode(self): + derive = SchedulerRecvSkipper.derive_forward_mode + decode = ForwardMode.DECODE.value + extend = ForwardMode.EXTEND.value + mixed = ForwardMode.MIXED.value + idle = ForwardMode.IDLE.value + prebuilt = ForwardMode.PREBUILT.value + verify = ForwardMode.TARGET_VERIFY.value + + # All ranks idle/prebuilt: same bucket as "no last batch". + self.assertIsNone(derive([idle, idle])) + self.assertIsNone(derive([prebuilt, idle])) + + # Any extend-like rank forces the immediate-recv bucket. + self.assertEqual(derive([decode, extend, idle]), ForwardMode.EXTEND) + self.assertEqual(derive([mixed, decode]), ForwardMode.EXTEND) + + # Pure decode-like steps keep the slow-recv weights. + self.assertEqual(derive([decode, idle, decode]), ForwardMode.DECODE) + self.assertEqual(derive([verify, verify]), ForwardMode.TARGET_VERIFY) + self.assertEqual(derive([verify, decode]), ForwardMode.DECODE) + + +if __name__ == "__main__": + unittest.main()