[Mamba] Fix checkpoint depth for prefixes that end off the radix page (#39115)
Co-authored-by: Ke Bao <ispobaoke@gmail.com>
This commit is contained in:
@@ -2965,7 +2965,18 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
# to force the math calculation to retrieve the correct mamba state from h.
|
# to force the math calculation to retrieve the correct mamba state from h.
|
||||||
return i + 1
|
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[
|
track_index = req.kv.mamba_ping_pong_track_buffer[
|
||||||
req.kv.mamba_next_track_idx
|
req.kv.mamba_next_track_idx
|
||||||
].item()
|
].item()
|
||||||
@@ -2978,14 +2989,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
# otherwise retrieved from h (i.e. unaligned).
|
# otherwise retrieved from h (i.e. unaligned).
|
||||||
# We need to pass the non-aligned seqlen to the calculation. Even though
|
# 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.
|
# 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 = seq_end
|
||||||
|
|
||||||
# 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
|
|
||||||
)
|
|
||||||
|
|
||||||
# A coarser checkpoint grid may not be a model-state boundary, so
|
# A coarser checkpoint grid may not be a model-state boundary, so
|
||||||
# force retrieval from the intermediate h state in that case.
|
# force retrieval from the intermediate h state in that case.
|
||||||
|
|||||||
@@ -3,7 +3,10 @@
|
|||||||
DCP widens the tree page past the mamba chunk grid. A checkpoint picked on the
|
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
|
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
|
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
|
import unittest
|
||||||
@@ -27,8 +30,9 @@ register_cpu_ci(est_time=10, suite="base-a-test-cpu")
|
|||||||
CHUNK = 64
|
CHUNK = 64
|
||||||
|
|
||||||
|
|
||||||
def _track_seqlen(*, tree_page: int, prefix_len: int, extend_len: int) -> int:
|
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."""
|
"""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)
|
server_args = ServerArgs(model_path="dummy", page_size=CHUNK)
|
||||||
# The property would otherwise load the HF config for the dummy model.
|
# The property would otherwise load the HF config for the dummy model.
|
||||||
server_args._mamba_cache_chunk_size = CHUNK
|
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 = MagicMock()
|
||||||
batch.req_to_token_pool.get_mamba_ping_pong_other_idx.return_value = 1
|
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
|
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)
|
depth = _track_seqlen(tree_page=CHUNK, prefix_len=16384, extend_len=4066)
|
||||||
self.assertEqual(depth, 20416)
|
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):
|
class TestMambaTrackGrid(unittest.TestCase):
|
||||||
def _grid(self, *, interval: int, tree_page: int, chunk: int = CHUNK) -> int:
|
def _grid(self, *, interval: int, tree_page: int, chunk: int = CHUNK) -> int:
|
||||||
|
|||||||
Reference in New Issue
Block a user