From 6f69f927da9e5692bb4709821faecff6be9b5a8a Mon Sep 17 00:00:00 2001 From: "Po-Han Huang (NVIDIA)" <53919306+nvpohanh@users.noreply.github.com> Date: Thu, 20 Aug 2026 03:01:36 +0800 Subject: [PATCH] [Scheduler] Add configurable decode interval after prefill (#35017) --- python/sglang/srt/managers/scheduler.py | 27 +++++++ python/sglang/srt/server_args.py | 10 +++ .../test_scheduler_chunked_req_gate.py | 2 + .../test_scheduler_prefill_decode_interval.py | 75 +++++++++++++++++++ .../unit/server_args/test_server_args.py | 9 +++ 5 files changed, 123 insertions(+) create mode 100644 test/registered/unit/managers/test_scheduler_prefill_decode_interval.py diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index ca3fb39df..b617ed301 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -1159,6 +1159,8 @@ class Scheduler( def init_chunked_prefill(self): self.chunked_prefill_size = get_schedule().chunked_prefill_size + self.prefill_decode_interval = get_schedule().prefill_decode_interval + self._prefill_decode_interval_remaining = 0 uses_transformers_backend = ( get_resolved_model_impl(self.model_config) == ModelImpl.TRANSFORMERS ) @@ -1195,6 +1197,28 @@ class Scheduler( ) self.enable_dynamic_chunking = False + def _should_defer_prefill(self) -> bool: + if self._prefill_decode_interval_remaining == 0: + return False + + self._prefill_decode_interval_remaining -= 1 + return True + + def _arm_prefill_decode_interval(self, batch: Optional[ScheduleBatch]) -> None: + if self.prefill_decode_interval == 0 or batch is None: + return + + # DP attention synchronizes this flag across ranks. This keeps every + # rank on the same prefill/decode cadence even when only one rank has + # local prefill work. Non-DP scheduling can use the local mode directly. + is_extend = ( + batch.is_extend_in_batch + if self.require_mlp_sync + else batch.forward_mode.is_extend() + ) + if is_extend: + self._prefill_decode_interval_remaining = self.prefill_decode_interval + def init_metrics_reporter( self, tp_rank: int, pp_rank: int, dp_rank: Optional[int] ) -> None: @@ -3128,6 +3152,8 @@ class Scheduler( if self.dllm_config is not None: new_batch = self.get_new_batch_dllm(running_batch) + elif self._should_defer_prefill(): + new_batch = None else: prefill_plan = self.get_new_batch_prefill(running_batch) new_batch = prefill_plan.batch_to_run @@ -3161,6 +3187,7 @@ class Scheduler( ret = self.dp_attn_adapter.maybe_prepare_mlp_sync_batch( ret, need_sync=need_mlp_sync ) + self._arm_prefill_decode_interval(ret) # Handle ngram embedding ret = self.ngram_embedding_manager.prepare_for_forward( diff --git a/python/sglang/srt/server_args.py b/python/sglang/srt/server_args.py index a64a523cf..7ea687fe9 100644 --- a/python/sglang/srt/server_args.py +++ b/python/sglang/srt/server_args.py @@ -804,6 +804,11 @@ class ServerArgs: "The maximum number of tokens in a chunk for the chunked prefill. Setting this to -1 means disabling chunked prefill.", NS("schedule"), ] = None + prefill_decode_interval: A[ + int, + "The number of decode rounds to run after a prefill batch before scheduling the next prefill. In data-parallel attention mode, the interval is synchronized across all DP ranks. Set to 0 to disable.", + NS("schedule"), + ] = 0 enable_dynamic_chunking: A[ bool, "Enable dynamic chunk size adjustment for pipeline parallelism. When enabled, chunk sizes are dynamically calculated based on fitted function to maintain consistent execution time across chunks.", @@ -3640,6 +3645,7 @@ class ServerArgs: self._handle_return_hidden_states_mode() self._handle_media_url_security() self._handle_hicache_ratio_default() + self._validate_prefill_decode_interval() if self.model_path.lower() in ["none", "dummy"]: return @@ -8690,6 +8696,10 @@ class ServerArgs: f"(got {self.asr_max_concurrent_sessions})." ) + def _validate_prefill_decode_interval(self): + if self.prefill_decode_interval < 0: + raise ValueError("--prefill-decode-interval must be non-negative.") + def _handle_other_validations(self): if self.default_chat_template_kwargs is not None and not isinstance( self.default_chat_template_kwargs, dict diff --git a/test/registered/unit/managers/test_scheduler_chunked_req_gate.py b/test/registered/unit/managers/test_scheduler_chunked_req_gate.py index 35676f187..4d15cdc94 100644 --- a/test/registered/unit/managers/test_scheduler_chunked_req_gate.py +++ b/test/registered/unit/managers/test_scheduler_chunked_req_gate.py @@ -91,6 +91,8 @@ def _scheduler_for_get_next_batch(*, tree_cache, chunked_req) -> Scheduler: s.running_batch.is_prefill_only = False s.running_batch.batch_is_full = False s.running_batch.reqs = [] + s.prefill_decode_interval = 0 + s._prefill_decode_interval_remaining = 0 s.get_new_batch_prefill = MagicMock( return_value=NextBatchPlan(batch_to_run=None, running_batch=s.running_batch) ) diff --git a/test/registered/unit/managers/test_scheduler_prefill_decode_interval.py b/test/registered/unit/managers/test_scheduler_prefill_decode_interval.py new file mode 100644 index 000000000..8b0d683c4 --- /dev/null +++ b/test/registered/unit/managers/test_scheduler_prefill_decode_interval.py @@ -0,0 +1,75 @@ +"""Tests for scheduler prefill/decode interleaving.""" + +import unittest +from types import SimpleNamespace + +from sglang.test.ci.ci_register import register_cpu_ci +from sglang.test.test_utils import maybe_stub_sgl_kernel + +maybe_stub_sgl_kernel() + +from sglang.srt.managers.scheduler import Scheduler + +register_cpu_ci(est_time=1, suite="base-a-test-cpu") + + +def _make_scheduler(*, interval: int, require_mlp_sync: bool) -> Scheduler: + scheduler = Scheduler.__new__(Scheduler) + scheduler.prefill_decode_interval = interval + scheduler._prefill_decode_interval_remaining = 0 + scheduler.require_mlp_sync = require_mlp_sync + return scheduler + + +def _make_batch(*, local_extend: bool, global_extend: bool): + return SimpleNamespace( + forward_mode=SimpleNamespace(is_extend=lambda: local_extend), + is_extend_in_batch=global_extend, + ) + + +class TestPrefillDecodeInterval(unittest.TestCase): + def test_disabled_interval_does_not_arm(self): + scheduler = _make_scheduler(interval=0, require_mlp_sync=False) + + scheduler._arm_prefill_decode_interval( + _make_batch(local_extend=True, global_extend=False) + ) + + self.assertFalse(scheduler._should_defer_prefill()) + + def test_non_dp_interval_uses_local_forward_mode(self): + scheduler = _make_scheduler(interval=2, require_mlp_sync=False) + + scheduler._arm_prefill_decode_interval( + _make_batch(local_extend=True, global_extend=False) + ) + + self.assertTrue(scheduler._should_defer_prefill()) + self.assertTrue(scheduler._should_defer_prefill()) + self.assertFalse(scheduler._should_defer_prefill()) + + def test_dp_interval_uses_globally_synchronized_extend_flag(self): + scheduler = _make_scheduler(interval=2, require_mlp_sync=True) + + # This rank is locally decoding, but another DP rank is prefilling. + scheduler._arm_prefill_decode_interval( + _make_batch(local_extend=False, global_extend=True) + ) + + self.assertEqual(scheduler._prefill_decode_interval_remaining, 2) + self.assertTrue(scheduler._should_defer_prefill()) + + def test_decode_batch_does_not_rearm_interval(self): + scheduler = _make_scheduler(interval=2, require_mlp_sync=True) + scheduler._prefill_decode_interval_remaining = 1 + + scheduler._arm_prefill_decode_interval( + _make_batch(local_extend=False, global_extend=False) + ) + + self.assertEqual(scheduler._prefill_decode_interval_remaining, 1) + + +if __name__ == "__main__": + unittest.main() diff --git a/test/registered/unit/server_args/test_server_args.py b/test/registered/unit/server_args/test_server_args.py index 6117b5560..c94adc5e7 100644 --- a/test/registered/unit/server_args/test_server_args.py +++ b/test/registered/unit/server_args/test_server_args.py @@ -43,6 +43,15 @@ _mock_device.start() class TestPrepareServerArgs(CustomTestCase): + def test_prefill_decode_interval(self): + args = ServerArgs(model_path="dummy", prefill_decode_interval=16) + self.assertEqual(args.prefill_decode_interval, 16) + + with self.assertRaisesRegex( + ValueError, "--prefill-decode-interval must be non-negative" + ): + ServerArgs(model_path="dummy", prefill_decode_interval=-1) + def test_return_hidden_states_mode_configuration(self): disabled = ServerArgs(model_path="dummy") self.assertFalse(disabled.enable_return_hidden_states)