diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index f1e7d5259..23de35caa 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -4,7 +4,6 @@ from sglang.srt.dllm.config import DllmConfig from sglang.srt.model_executor.forward_batch_info import ForwardBatch from sglang.srt.runtime_context import ( get_disagg, - get_exec, get_schedule, get_serving, get_spec, @@ -12,6 +11,7 @@ from sglang.srt.runtime_context import ( mamba_checkpoint_grid, mamba_extra_buffer_enabled, mamba_extra_buffer_lazy_enabled, + mamba_track_grid, ) from sglang.srt.utils.common import ( Range, @@ -3107,7 +3107,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): ) if mamba_extra_buffer_enabled(): - mamba_track_interval = get_exec().mamba.mamba_track_interval + mamba_track_interval = mamba_track_grid(self.tree_cache.page_size) if len(self.reqs) == 0: self.mamba_track_indices = torch.empty( diff --git a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py index 024781c37..93e37a010 100644 --- a/python/sglang/srt/managers/scheduler_components/batch_result_processor.py +++ b/python/sglang/srt/managers/scheduler_components/batch_result_processor.py @@ -34,10 +34,10 @@ from sglang.srt.model_executor.forward_batch_info import ( ) from sglang.srt.runtime_context import ( get_disagg, - get_exec, get_memory, get_observability, mamba_extra_buffer_lazy_enabled, + mamba_track_grid, max_speculative_num_draft_tokens, ) from sglang.srt.speculative.base_spec_worker import BaseSpecWorker @@ -1131,7 +1131,7 @@ class SchedulerBatchResultProcessor: if known_boundary: self._mamba_assert_committed_len_lookahead(req) track_seqlen = req.kv_committed_len - assert track_seqlen % get_exec().mamba.mamba_track_interval == 0 + assert track_seqlen % mamba_track_grid(self.tree_cache.page_size) == 0 at_boundary = True else: at_boundary, track_seqlen = self._mamba_check_track_boundary( @@ -1191,7 +1191,7 @@ class SchedulerBatchResultProcessor: other_idx ].item() == -1 and mamba_lazy_spec_in_window( req, - get_exec().mamba.mamba_track_interval, + mamba_track_grid(self.tree_cache.page_size), max_speculative_num_draft_tokens(), ) if ( @@ -1244,7 +1244,7 @@ class SchedulerBatchResultProcessor: For spec decode, the boundary is detected by comparing the accepted seq_len range against interval boundaries. """ - interval = get_exec().mamba.mamba_track_interval + interval = mamba_track_grid(self.tree_cache.page_size) if batch.spec_algorithm.is_none(): lookahead = req.decode_batch_idx - batch.mamba_decode_batch_idx_cpu[i] diff --git a/python/sglang/srt/runtime_context.py b/python/sglang/srt/runtime_context.py index 57b0c1051..f769853fe 100644 --- a/python/sglang/srt/runtime_context.py +++ b/python/sglang/srt/runtime_context.py @@ -1490,6 +1490,14 @@ def mamba_checkpoint_grid(tree_page: int) -> int: return math.lcm(mamba_cache_chunk_size(), tree_page) +def mamba_track_grid(tree_page: int) -> int: + """The same granularity for a decode-donated checkpoint, which additionally + has to land on the requested ``mamba_track_interval``.""" + return math.lcm( + mamba_checkpoint_grid(tree_page), get_exec().mamba.mamba_track_interval + ) + + def max_speculative_num_draft_tokens() -> int | None: """The largest draft-token count speculative decoding may use. diff --git a/python/sglang/srt/speculative/dflash_worker_v2.py b/python/sglang/srt/speculative/dflash_worker_v2.py index 7ae69921b..98bee8b5a 100644 --- a/python/sglang/srt/speculative/dflash_worker_v2.py +++ b/python/sglang/srt/speculative/dflash_worker_v2.py @@ -32,7 +32,7 @@ from sglang.srt.model_executor.forward_batch_info import ( ForwardMode, compute_position, ) -from sglang.srt.runtime_context import get_exec, get_schedule +from sglang.srt.runtime_context import get_exec, get_schedule, mamba_track_grid from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.base_spec_worker import BaseSpecWorker from sglang.srt.speculative.dflash_info import DFlashVerifyInput @@ -1526,7 +1526,7 @@ class DFlashWorkerV2(BaseSpecWorker): mamba_steps_to_track = None if batch.mamba_track_indices is not None: - mamba_track_interval = get_exec().mamba.mamba_track_interval + mamba_track_interval = mamba_track_grid(batch.tree_cache.page_size) to_track_mask = ( seq_lens_pre_verify // mamba_track_interval != batch.seq_lens // mamba_track_interval diff --git a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py index d3024c445..4788947a3 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py @@ -19,7 +19,13 @@ from sglang.srt.model_executor.forward_batch_info import ( CaptureHiddenMode, compute_position, ) -from sglang.srt.runtime_context import get_exec, get_parallel, get_schedule, get_spec +from sglang.srt.runtime_context import ( + get_exec, + get_parallel, + get_schedule, + get_spec, + mamba_track_grid, +) from sglang.srt.server_args import ServerArgs from sglang.srt.speculative.base_spec_worker import BaseSpecWorker from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2 @@ -815,7 +821,7 @@ class DSparkWorkerV2(BaseSpecWorker): mamba_steps_to_track = None if batch.mamba_track_indices is not None: - mamba_track_interval = get_exec().mamba.mamba_track_interval + mamba_track_interval = mamba_track_grid(batch.tree_cache.page_size) to_track_mask = ( seq_lens_pre_verify // mamba_track_interval != seq_lens_post_verify // mamba_track_interval diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index a2d6d91e2..b0fb39c53 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -47,10 +47,10 @@ from sglang.srt.mem_cache.allocation import ( assign_req_to_token_pool_func as assign_req_to_token_pool_func, ) from sglang.srt.runtime_context import ( - get_exec, get_spec, mamba_extra_buffer_enabled, mamba_extra_buffer_lazy_enabled, + mamba_track_grid, max_speculative_num_draft_tokens, ) from sglang.srt.utils import ( @@ -811,7 +811,7 @@ def _verify_commit_step_indices( return last_correct_step_indices, None seq_lens_pre_verify = batch.seq_lens seq_lens_post_verify = batch.seq_lens + accept_lens - mamba_track_interval = get_exec().mamba.mamba_track_interval + mamba_track_interval = mamba_track_grid(batch.tree_cache.page_size) to_track_mask = ( seq_lens_pre_verify // mamba_track_interval != seq_lens_post_verify // mamba_track_interval @@ -984,7 +984,7 @@ def commit_mamba_states_after_verify( mamba_track_indices = batch.mamba_track_indices mamba_steps_to_track = None if mamba_track_indices is not None: - ti = get_exec().mamba.mamba_track_interval + ti = mamba_track_grid(batch.tree_cache.page_size) seq_pre = batch.seq_lens seq_post = batch.seq_lens + accept_lens to_track_mask = seq_pre // ti != seq_post // ti @@ -1036,7 +1036,7 @@ def spec_prepare_for_decode(batch: ScheduleBatch) -> None: if mamba_extra_buffer_lazy_enabled(): # Scheduler phase (outside forward isolation). batch.mamba_lazy_spec_prepare( - get_exec().mamba.mamba_track_interval, + mamba_track_grid(batch.tree_cache.page_size), max_speculative_num_draft_tokens(), ) if batch.spec_algorithm.is_dflash_family(): diff --git a/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py b/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py index 17431c6c9..8173cd23d 100644 --- a/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py +++ b/test/registered/unit/managers/test_batch_result_processor_mamba_boundary.py @@ -17,6 +17,11 @@ from sglang.test.ci.ci_register import register_cpu_ci register_cpu_ci(est_time=1, suite="base-a-test-cpu") +# The decode checkpoint grid is lcm(mamba_cache_chunk_size, tree page, +# interval); keeping all three equal leaves it at the interval under test. +TRACK_INTERVAL = 4 + + def _make_batch() -> tuple[Req, ScheduleBatch]: sampling_params = SamplingParams(max_new_tokens=32) sampling_params.normalize(None) @@ -31,6 +36,7 @@ def _make_batch() -> tuple[Req, ScheduleBatch]: req.kv_committed_len = 2 batch = ScheduleBatch(reqs=[req]) + batch.tree_cache = SimpleNamespace(page_size=TRACK_INTERVAL) batch.device = "cpu" batch.model_config = SimpleNamespace(is_encoder_decoder=False) batch.enable_overlap = True @@ -57,7 +63,7 @@ def _make_processor() -> SchedulerBatchResultProcessor: server_args=SimpleNamespace(), model_config=SimpleNamespace(think_end_ids=None), token_to_kv_pool_allocator=MagicMock(), - tree_cache=None, + tree_cache=SimpleNamespace(page_size=TRACK_INTERVAL), hisparse_coordinator=None, req_to_token_pool=None, decode_offload_manager=None, @@ -154,7 +160,8 @@ class TestMambaBoundaryMaskReuse(unittest.TestCase): # defaults. get_context().override_server_args( mamba_radix_cache_strategy="extra_buffer", - mamba_track_interval=4, + mamba_track_interval=TRACK_INTERVAL, + _mamba_cache_chunk_size=TRACK_INTERVAL, ), patch( "sglang.srt.managers.schedule_batch.alloc_for_decode", diff --git a/test/registered/unit/managers/test_mamba_checkpoint_depth.py b/test/registered/unit/managers/test_mamba_checkpoint_depth.py index beb937949..ec8c3720a 100644 --- a/test/registered/unit/managers/test_mamba_checkpoint_depth.py +++ b/test/registered/unit/managers/test_mamba_checkpoint_depth.py @@ -14,6 +14,7 @@ from unittest.mock import MagicMock import torch from sglang.srt.managers.schedule_batch import Req, ScheduleBatch +from sglang.srt.runtime_context import get_context, mamba_track_grid from sglang.srt.sampling.sampling_params import SamplingParams from sglang.srt.server_args import ( ServerArgs, @@ -21,7 +22,7 @@ from sglang.srt.server_args import ( ) from sglang.test.ci.ci_register import register_cpu_ci -register_cpu_ci(est_time=1, suite="base-a-test-cpu") +register_cpu_ci(est_time=2, suite="base-a-test-cpu") CHUNK = 64 @@ -70,5 +71,24 @@ class TestMambaCheckpointDepth(unittest.TestCase): self.assertEqual(depth, 20416) +class TestMambaTrackGrid(unittest.TestCase): + def _grid(self, *, interval: int, tree_page: int, chunk: int = CHUNK) -> int: + with get_context().override_server_args( + mamba_track_interval=interval, _mamba_cache_chunk_size=chunk + ): + return mamba_track_grid(tree_page) + + def test_widened_tree_page_rounds_the_interval_up(self): + self.assertEqual(self._grid(interval=256, tree_page=512), 512) + + def test_interval_already_on_the_tree_page_is_untouched(self): + self.assertEqual(self._grid(interval=256, tree_page=256), 256) + self.assertEqual(self._grid(interval=256, tree_page=128), 256) + self.assertEqual(self._grid(interval=256, tree_page=64), 256) + + def test_grid_stays_on_the_chunk_size(self): + self.assertEqual(self._grid(interval=192, tree_page=64, chunk=128) % 128, 0) + + if __name__ == "__main__": unittest.main() diff --git a/test/registered/unit/spec/test_ngram_mamba_verify_update.py b/test/registered/unit/spec/test_ngram_mamba_verify_update.py index a43eb5d51..a620b5f2f 100644 --- a/test/registered/unit/spec/test_ngram_mamba_verify_update.py +++ b/test/registered/unit/spec/test_ngram_mamba_verify_update.py @@ -198,8 +198,8 @@ class TestNgramMambaVerifyUpdate(CustomTestCase): "sglang.srt.speculative.spec_utils.mambaish_config", return_value={"some": "config"}, ), patch( - "sglang.srt.speculative.spec_utils.get_exec", - return_value=MagicMock(mamba=MagicMock(mamba_track_interval=256)), + "sglang.srt.speculative.spec_utils.mamba_track_grid", + return_value=256, ): commit_mamba_states_after_verify( target_worker,