[Scheduler] Add configurable decode interval after prefill (#35017)
This commit is contained in:
@@ -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(
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user