[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:
Khoa Pham
2026-08-20 12:15:42 -07:00
committed by GitHub
co-authored by Claude Opus 5 Ke Bao
parent 308bc1228b
commit eac91ac362
9 changed files with 60 additions and 19 deletions
+2 -2
View File
@@ -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(
@@ -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]
+8
View File
@@ -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.
@@ -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
@@ -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
+4 -4
View File
@@ -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():