[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:
co-authored by
xbfs
Cursor
parent
3ffe72517f
commit
7074704c0c
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user