From 5e73c89b34d8e4ae7422b2a3cce2a85508cf1101 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Sun, 16 Aug 2026 16:01:26 -0700 Subject: [PATCH] [Spec] Simplify compute_spec_v2_logprobs signature and skip identity gathers (#35058) --- python/sglang/srt/layers/logprob_processor.py | 49 +++++++++++-------- .../srt/speculative/dflash_worker_v2.py | 10 ++-- .../dspark_components/dspark_worker_v2.py | 21 ++------ .../srt/speculative/eagle_worker_common.py | 5 +- .../sglang/srt/speculative/eagle_worker_v2.py | 1 - .../multi_layer_eagle_worker_v2.py | 1 - python/sglang/srt/speculative/ngram_worker.py | 10 ++-- 7 files changed, 40 insertions(+), 57 deletions(-) diff --git a/python/sglang/srt/layers/logprob_processor.py b/python/sglang/srt/layers/logprob_processor.py index 997695bd0..1536030ae 100644 --- a/python/sglang/srt/layers/logprob_processor.py +++ b/python/sglang/srt/layers/logprob_processor.py @@ -13,6 +13,7 @@ from sglang.srt.runtime_context import get_exec if TYPE_CHECKING: from sglang.srt.layers.logits_processor import LogitsMetadata, LogitsProcessorOutput from sglang.srt.layers.vocab_parallel_embedding import VocabParallelEmbedding + from sglang.srt.managers.schedule_batch import ScheduleBatch logger = logging.getLogger(__name__) @@ -348,25 +349,36 @@ def get_token_ids_logprobs_chunk( return next_split_pruned_len -def compute_spec_v2_logprobs( - batch, - logits_output, +def compute_spec_logprobs( + batch: ScheduleBatch, + logits_output: LogitsProcessorOutput, predict: torch.Tensor, - accept_index: torch.Tensor, - speculative_num_steps: int, + *, + accept_index: Optional[torch.Tensor] = None, + chain_stride: Optional[int] = None, ): - """Compute logprobs for accepted tokens after spec v2 verify sampling. + assert (accept_index is None) != ( + chain_stride is None + ), "pass exactly one of accept_index / chain_stride" - Gathers logits at accepted positions, applies log_softmax (temperature-scaled - if not greedy), and populates logits_output.next_token_logprobs (plus optional - top-k / token-ids logprobs) so they flow through copy_to_cpu(). - """ bs = len(batch.seq_lens) - max_accept = speculative_num_steps + 1 - device = predict.device + next_token_logits = logits_output.next_token_logits - flat_accept_idx = accept_index.long().reshape(-1) - gathered_logits = logits_output.next_token_logits[flat_accept_idx] + if accept_index is not None: + max_accept = accept_index.shape[1] + flat_accept_idx = accept_index.long().reshape(-1) + gathered_logits = next_token_logits[flat_accept_idx] + accepted_token_ids = predict[flat_accept_idx] + else: + max_accept = chain_stride + # Guards the layout contract the identity gather rests on: out token + # (b, j) must come from logits row b * stride + j. + assert next_token_logits.shape[0] == bs * max_accept, ( + f"chain layout expects {bs * max_accept} logits rows, got " + f"{next_token_logits.shape[0]}" + ) + gathered_logits = next_token_logits + accepted_token_ids = predict if batch.sampling_info.is_all_greedy or envs.SGLANG_RETURN_ORIGINAL_LOGPROB.get(): gathered_logprobs = torch.nn.functional.log_softmax(gathered_logits, dim=-1) @@ -381,12 +393,9 @@ def compute_spec_v2_logprobs( ) gathered_logprobs.clamp_(min=torch.finfo(gathered_logprobs.dtype).min) - accepted_token_ids = predict[flat_accept_idx] - token_logprobs = gathered_logprobs[ - torch.arange(bs * max_accept, device=device), - accepted_token_ids.long(), - ] - logits_output.next_token_logprobs = token_logprobs.reshape(bs, max_accept) + logits_output.next_token_logprobs = gathered_logprobs.gather( + 1, accepted_token_ids.long().view(-1, 1) + ).view(bs, max_accept) if batch.top_logprobs_nums and any(x > 0 for x in batch.top_logprobs_nums): top_logprobs_nums_expanded = [ diff --git a/python/sglang/srt/speculative/dflash_worker_v2.py b/python/sglang/srt/speculative/dflash_worker_v2.py index 62e7340dd..e831d14cc 100644 --- a/python/sglang/srt/speculative/dflash_worker_v2.py +++ b/python/sglang/srt/speculative/dflash_worker_v2.py @@ -17,7 +17,7 @@ from sglang.srt.configs.hybrid_arch import mambaish_config from sglang.srt.distributed import get_tp_group from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.environ import envs -from sglang.srt.layers.logprob_processor import compute_spec_v2_logprobs +from sglang.srt.layers.logprob_processor import compute_spec_logprobs from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.scheduler import GenerationBatchResult from sglang.srt.managers.tp_worker import TpModelWorker @@ -1897,15 +1897,11 @@ class DFlashWorkerV2(BaseSpecWorker): new_seq_lens = None if batch.return_logprob: - output_indices = torch.arange( - bs * block_size, dtype=torch.int64, device=device - ).view(bs, block_size) - compute_spec_v2_logprobs( + compute_spec_logprobs( batch, logits_output, out_tokens.reshape(-1), - output_indices, - block_size - 1, + chain_stride=block_size, ) if self._need_mamba_verify_commit: diff --git a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py index d2aff43ac..f5dc7c795 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_worker_v2.py @@ -11,7 +11,7 @@ from sglang.kernels.ops.attention.dsv4.unified_kv_kernels.env_gate import ( from sglang.srt.configs.hybrid_arch import mambaish_config from sglang.srt.distributed.parallel_state_wrapper import ParallelState from sglang.srt.environ import envs -from sglang.srt.layers.logprob_processor import compute_spec_v2_logprobs +from sglang.srt.layers.logprob_processor import compute_spec_logprobs from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.scheduler import GenerationBatchResult from sglang.srt.managers.tp_worker import TpModelWorker @@ -127,7 +127,6 @@ class DSparkWorkerV2(BaseSpecWorker): self.draft_model_runner = bundle.draft_model_runner self.draft_model = bundle.draft_model self._draft_sampler = None - self._linear_accept_index_cache = None # The mask token is input-only (it is embedded, never sampled), so its # bound is the embedding-table row count: the PADDED vocab when the @@ -406,19 +405,6 @@ class DSparkWorkerV2(BaseSpecWorker): def note_request_finished(self, *, rid: str, natural_stop: bool) -> None: self._observers.note_request_finished(rid=rid, natural_stop=natural_stop) - def _linear_accept_indices(self, bs: int) -> torch.Tensor: - num_indices = bs * self.verify_num_draft_tokens - if ( - self._linear_accept_index_cache is None - or self._linear_accept_index_cache.numel() < num_indices - ): - self._linear_accept_index_cache = torch.arange( - num_indices, dtype=torch.int64, device=self.device - ) - return self._linear_accept_index_cache[:num_indices].view( - bs, self.verify_num_draft_tokens - ) - def forward_batch_generation( self, batch: ScheduleBatch, @@ -718,12 +704,11 @@ class DSparkWorkerV2(BaseSpecWorker): draft_tokens=draft_tokens, ) if batch.return_logprob: - compute_spec_v2_logprobs( + compute_spec_logprobs( batch, logits_output, accept.out_tokens.reshape(-1), - self._linear_accept_indices(bs), - self.verify_num_draft_tokens - 1, + chain_stride=self.verify_num_draft_tokens, ) if on_publish is not None: diff --git a/python/sglang/srt/speculative/eagle_worker_common.py b/python/sglang/srt/speculative/eagle_worker_common.py index 8e0c01826..91ce8a147 100644 --- a/python/sglang/srt/speculative/eagle_worker_common.py +++ b/python/sglang/srt/speculative/eagle_worker_common.py @@ -8,7 +8,7 @@ from sglang.kernels.ops.speculative.cache_locs import ( assign_draft_cache_locs_contiguous, ) from sglang.kernels.ops.speculative.eagle import fill_bonus_tokens_func -from sglang.srt.layers.logprob_processor import compute_spec_v2_logprobs +from sglang.srt.layers.logprob_processor import compute_spec_logprobs from sglang.srt.managers.utils import GenerationBatchResult from sglang.srt.model_executor.forward_batch_info import ( CaptureHiddenMode, @@ -467,7 +467,6 @@ def run_eagle_verify( plan_stream: Any, plan_stream_ctx: Any, topk: int, - num_steps: int, num_draft_tokens: int, device: str, metadata_ready_pre_pad: bool, @@ -625,7 +624,7 @@ def run_eagle_verify( bonus_tokens = torch.empty((0,), device=device, dtype=torch.int32) if batch.return_logprob and not batch.forward_mode.is_idle(): - compute_spec_v2_logprobs(batch, logits_output, predict, accept_index, num_steps) + compute_spec_logprobs(batch, logits_output, predict, accept_index=accept_index) if finalize_tree_path and not batch.forward_mode.is_idle() and topk > 1: # topk == 1 needs nothing here: the accepted path is already the front diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 7d6ddc45a..12bda19f1 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -1503,7 +1503,6 @@ class EAGLEWorkerV2(BaseSpecWorker): plan_stream=self.plan_stream, plan_stream_ctx=self.plan_stream_ctx, topk=self.topk, - num_steps=self.speculative_num_steps, num_draft_tokens=self.speculative_num_draft_tokens, device=self.device, metadata_ready_pre_pad=False, diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index 5ecef3a43..3af1dbc60 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -1038,7 +1038,6 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker): plan_stream=self.plan_stream, plan_stream_ctx=self.plan_stream_ctx, topk=self.topk, - num_steps=self.speculative_num_steps, num_draft_tokens=self.speculative_num_draft_tokens, device=self.device, metadata_ready_pre_pad=False, diff --git a/python/sglang/srt/speculative/ngram_worker.py b/python/sglang/srt/speculative/ngram_worker.py index 7901b0021..39caaa791 100644 --- a/python/sglang/srt/speculative/ngram_worker.py +++ b/python/sglang/srt/speculative/ngram_worker.py @@ -9,7 +9,7 @@ from sglang.kernels.ops.speculative.cache_locs import ( assign_extend_cache_locs_func as assign_extend_cache_locs_func, ) from sglang.srt.distributed.parallel_state_wrapper import ParallelState -from sglang.srt.layers.logprob_processor import compute_spec_v2_logprobs +from sglang.srt.layers.logprob_processor import compute_spec_logprobs from sglang.srt.managers.schedule_batch import ScheduleBatch from sglang.srt.managers.scheduler import GenerationBatchResult from sglang.srt.managers.tp_worker import TpModelWorker @@ -480,15 +480,11 @@ class NGRAMWorker(BaseSpecWorker): self.token_to_kv_pool_allocator, ) if batch.return_logprob: - # The last arg is the accept_index row width minus 1. NGRAM's - # accept_index is (bs, draft_token_num) -- the tree depth is not - # bounded by spec_steps like EAGLE's (bs, spec_steps + 1). - compute_spec_v2_logprobs( + compute_spec_logprobs( batch, logits_output, predict, - accept_index, - self.draft_token_num - 1, + accept_index=accept_index, ) if on_publish is not None: