[Spec] Dedup post-verify mamba state commit into shared spec_utils helpers (#27966)

Co-authored-by: xbfs <xuebf1@lenovo.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Liangsheng Yin
2026-06-11 18:33:40 -07:00
committed by GitHub
co-authored by xbfs Cursor
parent 3ffe72517f
commit 7074704c0c
4 changed files with 119 additions and 152 deletions
@@ -12,10 +12,7 @@ from sglang.srt.layers.dp_attention import (
is_dp_attention_enabled,
)
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.managers.schedule_batch import (
ScheduleBatch,
set_mamba_track_indices_from_reqs,
)
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.mem_cache.common import (
alloc_paged_token_slots_extend,
alloc_token_slots,
@@ -35,6 +32,7 @@ from sglang.srt.speculative.eagle_utils import verify_tree_greedy_func
from sglang.srt.speculative.spec_utils import (
SIMULATE_ACC_LEN,
generate_simulated_accept_index,
prepare_mamba_track_for_verify,
)
from sglang.srt.speculative.triton_ops.cache_locs import (
assign_draft_cache_locs_contiguous as assign_draft_cache_locs_contiguous,
@@ -412,10 +410,7 @@ class EagleVerifyInputV2Mixin:
device=device,
)
if get_global_server_args().enable_mamba_extra_buffer():
set_mamba_track_indices_from_reqs(batch)
batch.mamba_track_mask = None
batch.mamba_track_seqlens = None
prepare_mamba_track_for_verify(batch)
# TBO's split_spec_info reads these; no-verify-sync leaves both None.
self.seq_lens_cpu = batch.seq_lens_cpu
@@ -66,6 +66,7 @@ from sglang.srt.speculative.eagle_utils import (
)
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_utils import (
commit_mamba_states_after_verify,
draft_tp_context,
generate_token_bitmask,
load_token_map,
@@ -1285,12 +1286,13 @@ class EAGLEWorkerV2(BaseSpecWorker):
new_seq_lens = batch.seq_lens + accept_lens
# Update mamba state for hybrid GDN models after verification
if (
self.target_worker.model_runner.hybrid_gdn_config is not None
or self.target_worker.model_runner.mamba2_config is not None
or self.target_worker.model_runner.hybrid_lightning_config is not None
):
self._mamba_verify_update(batch, accept_lens, accept_index, bs)
commit_mamba_states_after_verify(
self.target_worker,
batch,
accept_lens,
accept_index,
self.speculative_num_draft_tokens,
)
if not batch.forward_mode.is_idle():
accept_tokens = predict[accept_index]
@@ -1337,65 +1339,6 @@ class EAGLEWorkerV2(BaseSpecWorker):
extra_keep_alive_refs=[verify_forward_batch],
)
def _mamba_verify_update(
self,
batch: ScheduleBatch,
accept_lens: torch.Tensor,
accept_index: torch.Tensor,
bs: int,
):
"""Update mamba state for hybrid GDN models after verification."""
# `accept_lens` already includes the bonus token (drafts + 1 per req).
if not batch.forward_mode.is_idle() and accept_index.numel() > 0:
accept_indices_offset = torch.arange(
0,
bs * self.speculative_num_draft_tokens,
step=self.speculative_num_draft_tokens,
dtype=accept_lens.dtype,
device=accept_lens.device,
)
req_idx = torch.arange(bs, dtype=torch.int64, device=accept_lens.device)
# Per-req tree step of the last accepted node, i.e. the step whose
# mamba state to commit; reduces to accept_lens - 1 for topk == 1.
last_correct_step_indices = (
accept_index[req_idx, (accept_lens - 1).to(torch.int64)]
- accept_indices_offset
)
if batch.mamba_track_indices is not None:
# If after verify, the request's seq_lens has crossed a mamba track interval,
# we need to update the mamba state for the request at the crossing point.
seq_lens_pre_verify = batch.seq_lens
seq_lens_post_verify = batch.seq_lens + accept_lens
mamba_track_interval = self.server_args.mamba_track_interval
to_track_mask = (
seq_lens_pre_verify // mamba_track_interval
!= seq_lens_post_verify // mamba_track_interval
)
tracking_point = (
seq_lens_post_verify // mamba_track_interval * mamba_track_interval
)
to_track_ith = torch.clamp(
tracking_point - seq_lens_pre_verify - 1, min=0
).to(torch.int64)
candidate_track_steps = (
accept_index[req_idx, to_track_ith] - accept_indices_offset
)
mamba_steps_to_track = torch.where(
to_track_mask,
candidate_track_steps,
torch.full_like(candidate_track_steps, -1),
)
else:
mamba_steps_to_track = None
self.target_worker.model_runner.attn_backend.update_mamba_state_after_mtp_verify(
last_correct_step_indices=last_correct_step_indices,
mamba_track_indices=batch.mamba_track_indices,
mamba_steps_to_track=mamba_steps_to_track,
model=self.target_worker.model_runner.model,
)
def _finalize_accept_tree_path(
self,
batch: ScheduleBatch,
+12 -79
View File
@@ -6,21 +6,20 @@ import torch
from sgl_kernel.speculative import reconstruct_indices_from_tree_mask
from sglang.srt.layers.utils.logprob import compute_spec_v2_logprobs
from sglang.srt.managers.schedule_batch import (
ScheduleBatch,
set_mamba_track_indices_from_reqs,
)
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.scheduler import GenerationBatchResult
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.observability.req_time_stats import set_time_batch
from sglang.srt.server_args import ServerArgs, get_global_server_args
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.base_spec_worker import BaseDraftWorker, BaseSpecWorker
from sglang.srt.speculative.cpp_ngram.ngram_corpus import NgramCorpus
from sglang.srt.speculative.ngram_info import NgramVerifyInput
from sglang.srt.speculative.spec_utils import (
commit_mamba_states_after_verify,
generate_token_bitmask,
move_accept_tokens_to_target_kvcache,
prepare_mamba_track_for_verify,
record_stream_for_v2_verify,
)
from sglang.srt.speculative.triton_ops.cache_locs import (
@@ -328,15 +327,7 @@ class NGRAMWorker(BaseSpecWorker):
device=self.device,
)
# Mirror EagleVerifyInputV2Mixin.prepare_for_v2_verify: spec batches skip
# the prepare_for_decode refresh and filter/merge null these fields, so
# rebuild track indices from reqs before verify. Clearing the mask also
# keeps a stale extend-time mask from triggering in-forward tracking
# during TARGET_VERIFY; tracking is done in _mamba_verify_update instead.
if get_global_server_args().enable_mamba_extra_buffer():
set_mamba_track_indices_from_reqs(batch)
batch.mamba_track_mask = None
batch.mamba_track_seqlens = None
prepare_mamba_track_for_verify(batch)
batch.spec_info = NgramVerifyInput(
draft_token=draft_tokens,
@@ -374,65 +365,6 @@ class NGRAMWorker(BaseSpecWorker):
i += 1
self.ngram_corpus.batch_put(batch_tokens)
def _mamba_verify_update(
self,
batch: ScheduleBatch,
accept_lens: torch.Tensor,
accept_index: torch.Tensor,
bs: int,
) -> None:
"""Commit accepted speculative states for hybrid linear attention backends."""
attn_backend = self.target_worker.model_runner.attn_backend
if not hasattr(attn_backend, "update_mamba_state_after_mtp_verify"):
return
if batch.forward_mode.is_idle() or accept_index.numel() == 0:
return
accept_indices_offset = torch.arange(
0,
bs * self.draft_token_num,
step=self.draft_token_num,
dtype=accept_lens.dtype,
device=accept_lens.device,
)
req_idx = torch.arange(bs, dtype=torch.int64, device=accept_lens.device)
last_correct_step_indices = (
accept_index[req_idx, (accept_lens - 1).to(torch.int64)]
- accept_indices_offset
)
if batch.mamba_track_indices is not None:
seq_lens_pre_verify = batch.seq_lens
seq_lens_post_verify = batch.seq_lens + accept_lens
mamba_track_interval = self.server_args.mamba_track_interval
to_track_mask = (
seq_lens_pre_verify // mamba_track_interval
!= seq_lens_post_verify // mamba_track_interval
)
tracking_point = (
seq_lens_post_verify // mamba_track_interval * mamba_track_interval
)
to_track_ith = torch.clamp(
tracking_point - seq_lens_pre_verify - 1, min=0
).to(torch.int64)
candidate_track_steps = (
accept_index[req_idx, to_track_ith] - accept_indices_offset
)
mamba_steps_to_track = torch.where(
to_track_mask,
candidate_track_steps,
torch.full_like(candidate_track_steps, -1),
)
else:
mamba_steps_to_track = None
attn_backend.update_mamba_state_after_mtp_verify(
last_correct_step_indices=last_correct_step_indices,
mamba_track_indices=batch.mamba_track_indices,
mamba_steps_to_track=mamba_steps_to_track,
model=self.target_worker.model_runner.model,
)
def forward_batch_generation(
self, batch: ScheduleBatch, on_publish=None
) -> GenerationBatchResult:
@@ -499,12 +431,13 @@ class NGRAMWorker(BaseSpecWorker):
accept_index,
) = verify_input.sample(batch, logits_output, vocab_mask)
new_seq_lens = batch.seq_lens + accept_lens
if (
self.target_worker.model_runner.hybrid_gdn_config is not None
or self.target_worker.model_runner.mamba2_config is not None
or self.target_worker.model_runner.hybrid_lightning_config is not None
):
self._mamba_verify_update(batch, accept_lens, accept_index, bs)
commit_mamba_states_after_verify(
self.target_worker,
batch,
accept_lens,
accept_index,
self.draft_token_num,
)
accept_tokens = predict[accept_index].flatten()
next_token_ids = accept_tokens
@@ -14,6 +14,7 @@ from sglang.srt.distributed.parallel_state import (
patch_tensor_parallel_group,
)
from sglang.srt.environ import envs
from sglang.srt.managers.schedule_batch import set_mamba_track_indices_from_reqs
from sglang.srt.server_args import get_global_server_args
from sglang.srt.speculative.triton_ops.cache_locs import (
align_evict_mask_to_page_size as align_evict_mask_to_page_size,
@@ -56,6 +57,7 @@ _is_musa = is_musa()
if TYPE_CHECKING:
from sglang.srt.constrained.base_grammar_backend import BaseGrammarObject
from sglang.srt.managers.schedule_batch import Req, ScheduleBatch
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.mem_cache.allocator import BaseTokenToKVPoolAllocator
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.eagle_info import EagleVerifyInput
@@ -536,3 +538,97 @@ def move_accept_tokens_to_target_kvcache(
token_to_kv_pool_allocator.get_kvcache().move_kv_cache(
tgt_cache_loc, accept_out_cache_loc
)
def prepare_mamba_track_for_verify(batch: ScheduleBatch) -> None:
"""Rebuild mamba track indices from reqs before a TARGET_VERIFY forward.
Spec batches skip the refresh in prepare_for_decode, and filter/merge
null these fields, so they must be rebuilt right before verify. Clearing
the mask also keeps a stale extend-time mask from triggering in-forward
tracking during TARGET_VERIFY; tracking is done in
commit_mamba_states_after_verify instead.
"""
if not get_global_server_args().enable_mamba_extra_buffer():
return
set_mamba_track_indices_from_reqs(batch)
batch.mamba_track_mask = None
batch.mamba_track_seqlens = None
def commit_mamba_states_after_verify(
target_worker: TpModelWorker,
batch: ScheduleBatch,
accept_lens: torch.Tensor,
accept_index: torch.Tensor,
draft_token_num: int,
) -> None:
"""Commit accepted per-step mamba states into the persistent caches.
During TARGET_VERIFY, hybrid linear attention backends keep per-step
states in intermediate caches instead of advancing the persistent
conv/ssm caches. After acceptance, the state of each request's last
accepted step is committed back, plus the interval-crossing state used
for prefix-cache tracking (mamba extra_buffer mode).
No-op for models without mamba-style state or backends without the
commit hook.
"""
model_runner = target_worker.model_runner
if model_runner.mambaish_config is None:
return
attn_backend = model_runner.attn_backend
if not hasattr(attn_backend, "update_mamba_state_after_mtp_verify"):
return
bs = accept_lens.shape[0]
# `accept_lens` already includes the bonus token (drafts + 1 per req).
if not batch.forward_mode.is_idle() and accept_index.numel() > 0:
accept_indices_offset = torch.arange(
0,
bs * draft_token_num,
step=draft_token_num,
dtype=accept_lens.dtype,
device=accept_lens.device,
)
req_idx = torch.arange(bs, dtype=torch.int64, device=accept_lens.device)
# Per-req tree step of the last accepted node, i.e. the step whose
# mamba state to commit; reduces to accept_lens - 1 for topk == 1.
last_correct_step_indices = (
accept_index[req_idx, (accept_lens - 1).to(torch.int64)]
- accept_indices_offset
)
if batch.mamba_track_indices is not None:
# If after verify, the request's seq_lens has crossed a mamba track interval,
# we need to update the mamba state for the request at the crossing point.
seq_lens_pre_verify = batch.seq_lens
seq_lens_post_verify = batch.seq_lens + accept_lens
mamba_track_interval = get_global_server_args().mamba_track_interval
to_track_mask = (
seq_lens_pre_verify // mamba_track_interval
!= seq_lens_post_verify // mamba_track_interval
)
tracking_point = (
seq_lens_post_verify // mamba_track_interval * mamba_track_interval
)
to_track_ith = torch.clamp(
tracking_point - seq_lens_pre_verify - 1, min=0
).to(torch.int64)
candidate_track_steps = (
accept_index[req_idx, to_track_ith] - accept_indices_offset
)
mamba_steps_to_track = torch.where(
to_track_mask,
candidate_track_steps,
torch.full_like(candidate_track_steps, -1),
)
else:
mamba_steps_to_track = None
attn_backend.update_mamba_state_after_mtp_verify(
last_correct_step_indices=last_correct_step_indices,
mamba_track_indices=batch.mamba_track_indices,
mamba_steps_to_track=mamba_steps_to_track,
model=model_runner.model,
)