diff --git a/python/sglang/srt/speculative/eagle_info_v2.py b/python/sglang/srt/speculative/eagle_info_v2.py index abcc63fad..1849480ba 100644 --- a/python/sglang/srt/speculative/eagle_info_v2.py +++ b/python/sglang/srt/speculative/eagle_info_v2.py @@ -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 diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 732a364bc..d3e4a3c61 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -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, diff --git a/python/sglang/srt/speculative/ngram_worker.py b/python/sglang/srt/speculative/ngram_worker.py index 306de76cf..0fef94b2c 100644 --- a/python/sglang/srt/speculative/ngram_worker.py +++ b/python/sglang/srt/speculative/ngram_worker.py @@ -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 diff --git a/python/sglang/srt/speculative/spec_utils.py b/python/sglang/srt/speculative/spec_utils.py index 8a7a88cb2..b94b5b3b9 100644 --- a/python/sglang/srt/speculative/spec_utils.py +++ b/python/sglang/srt/speculative/spec_utils.py @@ -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, + )