[Scheduler] Add configurable decode interval after prefill (#35017)

This commit is contained in:
Po-Han Huang (NVIDIA)
2026-08-19 12:01:36 -07:00
committed by GitHub
parent 4f8ecf6ae9
commit 6f69f927da
5 changed files with 123 additions and 0 deletions
+27
View File
@@ -1159,6 +1159,8 @@ class Scheduler(
def init_chunked_prefill(self): def init_chunked_prefill(self):
self.chunked_prefill_size = get_schedule().chunked_prefill_size 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 = ( uses_transformers_backend = (
get_resolved_model_impl(self.model_config) == ModelImpl.TRANSFORMERS get_resolved_model_impl(self.model_config) == ModelImpl.TRANSFORMERS
) )
@@ -1195,6 +1197,28 @@ class Scheduler(
) )
self.enable_dynamic_chunking = False 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( def init_metrics_reporter(
self, tp_rank: int, pp_rank: int, dp_rank: Optional[int] self, tp_rank: int, pp_rank: int, dp_rank: Optional[int]
) -> None: ) -> None:
@@ -3128,6 +3152,8 @@ class Scheduler(
if self.dllm_config is not None: if self.dllm_config is not None:
new_batch = self.get_new_batch_dllm(running_batch) new_batch = self.get_new_batch_dllm(running_batch)
elif self._should_defer_prefill():
new_batch = None
else: else:
prefill_plan = self.get_new_batch_prefill(running_batch) prefill_plan = self.get_new_batch_prefill(running_batch)
new_batch = prefill_plan.batch_to_run new_batch = prefill_plan.batch_to_run
@@ -3161,6 +3187,7 @@ class Scheduler(
ret = self.dp_attn_adapter.maybe_prepare_mlp_sync_batch( ret = self.dp_attn_adapter.maybe_prepare_mlp_sync_batch(
ret, need_sync=need_mlp_sync ret, need_sync=need_mlp_sync
) )
self._arm_prefill_decode_interval(ret)
# Handle ngram embedding # Handle ngram embedding
ret = self.ngram_embedding_manager.prepare_for_forward( ret = self.ngram_embedding_manager.prepare_for_forward(
+10
View File
@@ -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.", "The maximum number of tokens in a chunk for the chunked prefill. Setting this to -1 means disabling chunked prefill.",
NS("schedule"), NS("schedule"),
] = None ] = 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[ enable_dynamic_chunking: A[
bool, 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.", "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_return_hidden_states_mode()
self._handle_media_url_security() self._handle_media_url_security()
self._handle_hicache_ratio_default() self._handle_hicache_ratio_default()
self._validate_prefill_decode_interval()
if self.model_path.lower() in ["none", "dummy"]: if self.model_path.lower() in ["none", "dummy"]:
return return
@@ -8690,6 +8696,10 @@ class ServerArgs:
f"(got {self.asr_max_concurrent_sessions})." 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): def _handle_other_validations(self):
if self.default_chat_template_kwargs is not None and not isinstance( if self.default_chat_template_kwargs is not None and not isinstance(
self.default_chat_template_kwargs, dict self.default_chat_template_kwargs, dict
@@ -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.is_prefill_only = False
s.running_batch.batch_is_full = False s.running_batch.batch_is_full = False
s.running_batch.reqs = [] s.running_batch.reqs = []
s.prefill_decode_interval = 0
s._prefill_decode_interval_remaining = 0
s.get_new_batch_prefill = MagicMock( s.get_new_batch_prefill = MagicMock(
return_value=NextBatchPlan(batch_to_run=None, running_batch=s.running_batch) return_value=NextBatchPlan(batch_to_run=None, running_batch=s.running_batch)
) )
@@ -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()
@@ -43,6 +43,15 @@ _mock_device.start()
class TestPrepareServerArgs(CustomTestCase): 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): def test_return_hidden_states_mode_configuration(self):
disabled = ServerArgs(model_path="dummy") disabled = ServerArgs(model_path="dummy")
self.assertFalse(disabled.enable_return_hidden_states) self.assertFalse(disabled.enable_return_hidden_states)