diff --git a/python/sglang/srt/managers/schedule_batch.py b/python/sglang/srt/managers/schedule_batch.py index 5297d9a84..544768998 100755 --- a/python/sglang/srt/managers/schedule_batch.py +++ b/python/sglang/srt/managers/schedule_batch.py @@ -2965,7 +2965,18 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): # to force the math calculation to retrieve the correct mamba state from h. return i + 1 - mask = req.extend_range.length >= checkpoint_grid + # Pick the depth on the absolute checkpoint grid: a chunk boundary can leave + # the prefix off the (DCP-widened) tree page, and a prefix-relative depth then + # names a position no page can hold. Donate only where an h snapshot exists. + prefix_len = len(req.prefix_indices) + seq_end = prefix_len + req.extend_range.length + # mamba_track_seqlen_aligned/mamba_last_track_seqlen is actual tracked seqlen. Used to pass to + # mamba radix cache to track which seqlen this mamba state should store at. + mamba_track_seqlen_aligned = (seq_end // checkpoint_grid) * checkpoint_grid + mask = ( + mamba_track_seqlen_aligned > prefix_len + and (mamba_track_seqlen_aligned - prefix_len) % cache_chunk_size == 0 + ) track_index = req.kv.mamba_ping_pong_track_buffer[ req.kv.mamba_next_track_idx ].item() @@ -2978,14 +2989,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin): # otherwise retrieved from h (i.e. unaligned). # We need to pass the non-aligned seqlen to the calculation. Even though # we pass in mamba_track_seqlen, the actual tracked seqlen is mamba_last_track_seqlen. - mamba_track_seqlen = len(req.prefix_indices) + req.extend_range.length - - # mamba_track_seqlen_aligned/mamba_last_track_seqlen is actual tracked seqlen. Used to pass to - # mamba radix cache to track which seqlen this mamba state should store at. - mamba_track_seqlen_aligned = ( - len(req.prefix_indices) - + (req.extend_range.length // checkpoint_grid) * checkpoint_grid - ) + mamba_track_seqlen = seq_end # A coarser checkpoint grid may not be a model-state boundary, so # force retrieval from the intermediate h state in that case. diff --git a/test/registered/unit/managers/test_mamba_checkpoint_depth.py b/test/registered/unit/managers/test_mamba_checkpoint_depth.py index a08843cf6..dd4e62c03 100644 --- a/test/registered/unit/managers/test_mamba_checkpoint_depth.py +++ b/test/registered/unit/managers/test_mamba_checkpoint_depth.py @@ -3,7 +3,10 @@ DCP widens the tree page past the mamba chunk grid. A checkpoint picked on the finer grid names a depth no radix node can carry, so it gets attached to the preceding node and a later request resumes from a state that already covers -tokens past that node. +tokens past that node. The same happens when the depth is measured from a +prefix that a chunked prefill left on the scheduler page but off the widened +tree page: 64 + 512 = 576 is no 512-page boundary, and the finished request +then frees two kv-row segments that share the page holding 576. """ import unittest @@ -27,8 +30,9 @@ register_cpu_ci(est_time=10, suite="base-a-test-cpu") CHUNK = 64 -def _track_seqlen(*, tree_page: int, prefix_len: int, extend_len: int) -> int: - """Run one extend through the tracker and report the donated depth.""" +def _track_seqlen(*, tree_page: int, prefix_len: int, extend_len: int) -> int | None: + """Run one extend through the tracker and report the donated depth, or + None when the extend donates no checkpoint.""" server_args = ServerArgs(model_path="dummy", page_size=CHUNK) # The property would otherwise load the HF config for the dummy model. server_args._mamba_cache_chunk_size = CHUNK @@ -57,7 +61,10 @@ def _track_seqlen(*, tree_page: int, prefix_len: int, extend_len: int) -> int: batch.req_to_token_pool = MagicMock() batch.req_to_token_pool.get_mamba_ping_pong_other_idx.return_value = 1 - batch._mamba_radix_cache_v2_req_prepare_for_extend(req) + entry = batch._mamba_radix_cache_v2_req_prepare_for_extend(req) + if not entry.track_mask: + assert req.kv.mamba_last_track_seqlen is None + return None return req.kv.mamba_last_track_seqlen @@ -73,6 +80,25 @@ class TestMambaCheckpointDepth(unittest.TestCase): depth = _track_seqlen(tree_page=CHUNK, prefix_len=16384, extend_len=4066) self.assertEqual(depth, 20416) + def test_depth_is_picked_on_the_absolute_grid(self): + # (tree_page, prefix_len, extend_len) -> donated depth, or None when the + # extend crosses no page or the prefix is off the kernel chunk grid. + for tree_page, prefix_len, extend_len, expected in ( + (512, 64, 540, 512), + (512, 448, 128, 512), + (512, 64, 400, None), + (512, 37, 540, None), + ): + with self.subTest(prefix_len=prefix_len, extend_len=extend_len): + self.assertEqual( + _track_seqlen( + tree_page=tree_page, + prefix_len=prefix_len, + extend_len=extend_len, + ), + expected, + ) + class TestMambaTrackGrid(unittest.TestCase): def _grid(self, *, interval: int, tree_page: int, chunk: int = CHUNK) -> int: