[Fix] Land the decode mamba checkpoint depth on the tree page under DCP (#35412)
Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com> Co-authored-by: Ke Bao <ispobaoke@gmail.com>
This commit is contained in:
co-authored by
Claude Opus 5
Ke Bao
parent
308bc1228b
commit
eac91ac362
@@ -4,7 +4,6 @@ from sglang.srt.dllm.config import DllmConfig
|
|||||||
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
from sglang.srt.model_executor.forward_batch_info import ForwardBatch
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
get_disagg,
|
get_disagg,
|
||||||
get_exec,
|
|
||||||
get_schedule,
|
get_schedule,
|
||||||
get_serving,
|
get_serving,
|
||||||
get_spec,
|
get_spec,
|
||||||
@@ -12,6 +11,7 @@ from sglang.srt.runtime_context import (
|
|||||||
mamba_checkpoint_grid,
|
mamba_checkpoint_grid,
|
||||||
mamba_extra_buffer_enabled,
|
mamba_extra_buffer_enabled,
|
||||||
mamba_extra_buffer_lazy_enabled,
|
mamba_extra_buffer_lazy_enabled,
|
||||||
|
mamba_track_grid,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils.common import (
|
from sglang.srt.utils.common import (
|
||||||
Range,
|
Range,
|
||||||
@@ -3107,7 +3107,7 @@ class ScheduleBatch(ScheduleBatchDisaggregationDecodeMixin):
|
|||||||
)
|
)
|
||||||
|
|
||||||
if mamba_extra_buffer_enabled():
|
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:
|
if len(self.reqs) == 0:
|
||||||
self.mamba_track_indices = torch.empty(
|
self.mamba_track_indices = torch.empty(
|
||||||
|
|||||||
@@ -34,10 +34,10 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
get_disagg,
|
get_disagg,
|
||||||
get_exec,
|
|
||||||
get_memory,
|
get_memory,
|
||||||
get_observability,
|
get_observability,
|
||||||
mamba_extra_buffer_lazy_enabled,
|
mamba_extra_buffer_lazy_enabled,
|
||||||
|
mamba_track_grid,
|
||||||
max_speculative_num_draft_tokens,
|
max_speculative_num_draft_tokens,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
|
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
|
||||||
@@ -1131,7 +1131,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
if known_boundary:
|
if known_boundary:
|
||||||
self._mamba_assert_committed_len_lookahead(req)
|
self._mamba_assert_committed_len_lookahead(req)
|
||||||
track_seqlen = req.kv_committed_len
|
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
|
at_boundary = True
|
||||||
else:
|
else:
|
||||||
at_boundary, track_seqlen = self._mamba_check_track_boundary(
|
at_boundary, track_seqlen = self._mamba_check_track_boundary(
|
||||||
@@ -1191,7 +1191,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
other_idx
|
other_idx
|
||||||
].item() == -1 and mamba_lazy_spec_in_window(
|
].item() == -1 and mamba_lazy_spec_in_window(
|
||||||
req,
|
req,
|
||||||
get_exec().mamba.mamba_track_interval,
|
mamba_track_grid(self.tree_cache.page_size),
|
||||||
max_speculative_num_draft_tokens(),
|
max_speculative_num_draft_tokens(),
|
||||||
)
|
)
|
||||||
if (
|
if (
|
||||||
@@ -1244,7 +1244,7 @@ class SchedulerBatchResultProcessor:
|
|||||||
For spec decode, the boundary is detected by comparing the
|
For spec decode, the boundary is detected by comparing the
|
||||||
accepted seq_len range against interval boundaries.
|
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():
|
if batch.spec_algorithm.is_none():
|
||||||
lookahead = req.decode_batch_idx - batch.mamba_decode_batch_idx_cpu[i]
|
lookahead = req.decode_batch_idx - batch.mamba_decode_batch_idx_cpu[i]
|
||||||
|
|||||||
@@ -1490,6 +1490,14 @@ def mamba_checkpoint_grid(tree_page: int) -> int:
|
|||||||
return math.lcm(mamba_cache_chunk_size(), tree_page)
|
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:
|
def max_speculative_num_draft_tokens() -> int | None:
|
||||||
"""The largest draft-token count speculative decoding may use.
|
"""The largest draft-token count speculative decoding may use.
|
||||||
|
|
||||||
|
|||||||
@@ -32,7 +32,7 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
ForwardMode,
|
ForwardMode,
|
||||||
compute_position,
|
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.server_args import ServerArgs
|
||||||
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
|
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
|
||||||
from sglang.srt.speculative.dflash_info import DFlashVerifyInput
|
from sglang.srt.speculative.dflash_info import DFlashVerifyInput
|
||||||
@@ -1526,7 +1526,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
mamba_steps_to_track = None
|
mamba_steps_to_track = None
|
||||||
|
|
||||||
if batch.mamba_track_indices is not 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 = (
|
to_track_mask = (
|
||||||
seq_lens_pre_verify // mamba_track_interval
|
seq_lens_pre_verify // mamba_track_interval
|
||||||
!= batch.seq_lens // mamba_track_interval
|
!= batch.seq_lens // mamba_track_interval
|
||||||
|
|||||||
@@ -19,7 +19,13 @@ from sglang.srt.model_executor.forward_batch_info import (
|
|||||||
CaptureHiddenMode,
|
CaptureHiddenMode,
|
||||||
compute_position,
|
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.server_args import ServerArgs
|
||||||
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
|
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker
|
||||||
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
|
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
|
||||||
@@ -815,7 +821,7 @@ class DSparkWorkerV2(BaseSpecWorker):
|
|||||||
mamba_steps_to_track = None
|
mamba_steps_to_track = None
|
||||||
|
|
||||||
if batch.mamba_track_indices is not 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 = (
|
to_track_mask = (
|
||||||
seq_lens_pre_verify // mamba_track_interval
|
seq_lens_pre_verify // mamba_track_interval
|
||||||
!= seq_lens_post_verify // mamba_track_interval
|
!= seq_lens_post_verify // mamba_track_interval
|
||||||
|
|||||||
@@ -47,10 +47,10 @@ from sglang.srt.mem_cache.allocation import (
|
|||||||
assign_req_to_token_pool_func as assign_req_to_token_pool_func,
|
assign_req_to_token_pool_func as assign_req_to_token_pool_func,
|
||||||
)
|
)
|
||||||
from sglang.srt.runtime_context import (
|
from sglang.srt.runtime_context import (
|
||||||
get_exec,
|
|
||||||
get_spec,
|
get_spec,
|
||||||
mamba_extra_buffer_enabled,
|
mamba_extra_buffer_enabled,
|
||||||
mamba_extra_buffer_lazy_enabled,
|
mamba_extra_buffer_lazy_enabled,
|
||||||
|
mamba_track_grid,
|
||||||
max_speculative_num_draft_tokens,
|
max_speculative_num_draft_tokens,
|
||||||
)
|
)
|
||||||
from sglang.srt.utils import (
|
from sglang.srt.utils import (
|
||||||
@@ -811,7 +811,7 @@ def _verify_commit_step_indices(
|
|||||||
return last_correct_step_indices, None
|
return last_correct_step_indices, None
|
||||||
seq_lens_pre_verify = batch.seq_lens
|
seq_lens_pre_verify = batch.seq_lens
|
||||||
seq_lens_post_verify = batch.seq_lens + accept_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 = (
|
to_track_mask = (
|
||||||
seq_lens_pre_verify // mamba_track_interval
|
seq_lens_pre_verify // mamba_track_interval
|
||||||
!= seq_lens_post_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_track_indices = batch.mamba_track_indices
|
||||||
mamba_steps_to_track = None
|
mamba_steps_to_track = None
|
||||||
if mamba_track_indices is not 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_pre = batch.seq_lens
|
||||||
seq_post = batch.seq_lens + accept_lens
|
seq_post = batch.seq_lens + accept_lens
|
||||||
to_track_mask = seq_pre // ti != seq_post // ti
|
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():
|
if mamba_extra_buffer_lazy_enabled():
|
||||||
# Scheduler phase (outside forward isolation).
|
# Scheduler phase (outside forward isolation).
|
||||||
batch.mamba_lazy_spec_prepare(
|
batch.mamba_lazy_spec_prepare(
|
||||||
get_exec().mamba.mamba_track_interval,
|
mamba_track_grid(batch.tree_cache.page_size),
|
||||||
max_speculative_num_draft_tokens(),
|
max_speculative_num_draft_tokens(),
|
||||||
)
|
)
|
||||||
if batch.spec_algorithm.is_dflash_family():
|
if batch.spec_algorithm.is_dflash_family():
|
||||||
|
|||||||
@@ -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")
|
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]:
|
def _make_batch() -> tuple[Req, ScheduleBatch]:
|
||||||
sampling_params = SamplingParams(max_new_tokens=32)
|
sampling_params = SamplingParams(max_new_tokens=32)
|
||||||
sampling_params.normalize(None)
|
sampling_params.normalize(None)
|
||||||
@@ -31,6 +36,7 @@ def _make_batch() -> tuple[Req, ScheduleBatch]:
|
|||||||
req.kv_committed_len = 2
|
req.kv_committed_len = 2
|
||||||
|
|
||||||
batch = ScheduleBatch(reqs=[req])
|
batch = ScheduleBatch(reqs=[req])
|
||||||
|
batch.tree_cache = SimpleNamespace(page_size=TRACK_INTERVAL)
|
||||||
batch.device = "cpu"
|
batch.device = "cpu"
|
||||||
batch.model_config = SimpleNamespace(is_encoder_decoder=False)
|
batch.model_config = SimpleNamespace(is_encoder_decoder=False)
|
||||||
batch.enable_overlap = True
|
batch.enable_overlap = True
|
||||||
@@ -57,7 +63,7 @@ def _make_processor() -> SchedulerBatchResultProcessor:
|
|||||||
server_args=SimpleNamespace(),
|
server_args=SimpleNamespace(),
|
||||||
model_config=SimpleNamespace(think_end_ids=None),
|
model_config=SimpleNamespace(think_end_ids=None),
|
||||||
token_to_kv_pool_allocator=MagicMock(),
|
token_to_kv_pool_allocator=MagicMock(),
|
||||||
tree_cache=None,
|
tree_cache=SimpleNamespace(page_size=TRACK_INTERVAL),
|
||||||
hisparse_coordinator=None,
|
hisparse_coordinator=None,
|
||||||
req_to_token_pool=None,
|
req_to_token_pool=None,
|
||||||
decode_offload_manager=None,
|
decode_offload_manager=None,
|
||||||
@@ -154,7 +160,8 @@ class TestMambaBoundaryMaskReuse(unittest.TestCase):
|
|||||||
# defaults.
|
# defaults.
|
||||||
get_context().override_server_args(
|
get_context().override_server_args(
|
||||||
mamba_radix_cache_strategy="extra_buffer",
|
mamba_radix_cache_strategy="extra_buffer",
|
||||||
mamba_track_interval=4,
|
mamba_track_interval=TRACK_INTERVAL,
|
||||||
|
_mamba_cache_chunk_size=TRACK_INTERVAL,
|
||||||
),
|
),
|
||||||
patch(
|
patch(
|
||||||
"sglang.srt.managers.schedule_batch.alloc_for_decode",
|
"sglang.srt.managers.schedule_batch.alloc_for_decode",
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ from unittest.mock import MagicMock
|
|||||||
import torch
|
import torch
|
||||||
|
|
||||||
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
|
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.sampling.sampling_params import SamplingParams
|
||||||
from sglang.srt.server_args import (
|
from sglang.srt.server_args import (
|
||||||
ServerArgs,
|
ServerArgs,
|
||||||
@@ -21,7 +22,7 @@ from sglang.srt.server_args import (
|
|||||||
)
|
)
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
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
|
CHUNK = 64
|
||||||
|
|
||||||
@@ -70,5 +71,24 @@ class TestMambaCheckpointDepth(unittest.TestCase):
|
|||||||
self.assertEqual(depth, 20416)
|
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__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
@@ -198,8 +198,8 @@ class TestNgramMambaVerifyUpdate(CustomTestCase):
|
|||||||
"sglang.srt.speculative.spec_utils.mambaish_config",
|
"sglang.srt.speculative.spec_utils.mambaish_config",
|
||||||
return_value={"some": "config"},
|
return_value={"some": "config"},
|
||||||
), patch(
|
), patch(
|
||||||
"sglang.srt.speculative.spec_utils.get_exec",
|
"sglang.srt.speculative.spec_utils.mamba_track_grid",
|
||||||
return_value=MagicMock(mamba=MagicMock(mamba_track_interval=256)),
|
return_value=256,
|
||||||
):
|
):
|
||||||
commit_mamba_states_after_verify(
|
commit_mamba_states_after_verify(
|
||||||
target_worker,
|
target_worker,
|
||||||
|
|||||||
Reference in New Issue
Block a user