diff --git a/python/sglang/srt/model_executor/forward_batch_info.py b/python/sglang/srt/model_executor/forward_batch_info.py index 07a9e11d6..caf45e74e 100644 --- a/python/sglang/srt/model_executor/forward_batch_info.py +++ b/python/sglang/srt/model_executor/forward_batch_info.py @@ -760,9 +760,16 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): # process global_num_tokens and global_num_tokens_for_logprob if batch.spec_info is not None: - spec_info: SpecInput = batch.spec_info + from sglang.srt.speculative.spec_info import ( + spec_scale_global_num_tokens, + ) + global_num_tokens, global_num_tokens_for_logprob = ( - spec_info.get_spec_adjusted_global_num_tokens(batch) + spec_scale_global_num_tokens( + batch.spec_info, + batch.global_num_tokens, + batch.global_num_tokens_for_logprob, + ) ) else: global_num_tokens = batch.global_num_tokens @@ -1233,6 +1240,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): if self.forward_mode.is_idle(): self._original_forward_mode = self.forward_mode self.forward_mode = ForwardMode.TARGET_VERIFY + # Invert the spec_scale_global_num_tokens scaling. bs = self.batch_size = num_tokens // self.spec_info.num_tokens_per_req elif self.is_extend_in_batch and dp_padding_mode.is_max_len(): self._original_forward_mode = self.forward_mode @@ -1294,6 +1302,7 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin): self.extend_logprob_start_lens_cpu = self.extend_prefix_lens_cpu else: if self.spec_info is not None: + # Invert the spec_scale_global_num_tokens scaling. bs = self.batch_size = ( num_tokens // self.spec_info.num_tokens_per_req ) diff --git a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py index 0f28aa29b..1cde0bfd1 100644 --- a/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py +++ b/python/sglang/srt/model_executor/runner/decode_cuda_graph_runner.py @@ -511,13 +511,9 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner): return False if self.require_mlp_tp_gather: - cuda_graph_bs = ( - max(forward_batch.global_num_tokens_cpu) // self.num_tokens_per_req - if self.model_runner.spec_algorithm.is_eagle() - or self.model_runner.spec_algorithm.is_standalone() - or self.model_runner.spec_algorithm.is_dflash_family() - else max(forward_batch.global_num_tokens_cpu) - ) + # Raw sync values are per-rank request counts on decode-family + # rounds -- no width division, no per-algorithm enumeration. + cuda_graph_bs = max(forward_batch.original_global_num_tokens_cpu) else: cuda_graph_bs = forward_batch.batch_size diff --git a/python/sglang/srt/models/mindspore.py b/python/sglang/srt/models/mindspore.py index 25056adf9..5ad75b7f5 100644 --- a/python/sglang/srt/models/mindspore.py +++ b/python/sglang/srt/models/mindspore.py @@ -267,13 +267,13 @@ class MindSporeForCausalLM(torch.nn.Module): is_prefill = self._is_prefill(forward_batch) batch_valid_length = forward_batch.seq_lens.cpu().numpy() if forward_batch.forward_mode.is_target_verify(): - batch_valid_length += forward_batch.spec_info.num_tokens_per_req + batch_valid_length += forward_batch.spec_info.draft_token_num if forward_batch.extend_seq_lens is not None: q_seq_lens = forward_batch.extend_seq_lens.cpu().numpy() else: q_seq_lens = np.ones([forward_batch.batch_size], dtype=np.int32) if forward_batch.forward_mode.is_target_verify(): - q_seq_lens = q_seq_lens * forward_batch.spec_info.num_tokens_per_req + q_seq_lens = q_seq_lens * forward_batch.spec_info.draft_token_num page_size = get_token_to_kv_pool().page_size block_tables = tensor_torch2ms( diff --git a/python/sglang/srt/speculative/dspark_components/dspark_draft.py b/python/sglang/srt/speculative/dspark_components/dspark_draft.py index 7eb9a13f3..c4cb2e013 100644 --- a/python/sglang/srt/speculative/dspark_components/dspark_draft.py +++ b/python/sglang/srt/speculative/dspark_components/dspark_draft.py @@ -21,7 +21,10 @@ from sglang.srt.speculative.dspark_components.dspark_planner import VerifyWindow from sglang.srt.speculative.dspark_components.kernels.dspark_draft_model import ( SampleStepTokens, ) -from sglang.srt.speculative.spec_info import SpeculativeAlgorithm +from sglang.srt.speculative.spec_info import ( + SpeculativeAlgorithm, + spec_scale_global_num_tokens, +) from sglang.srt.speculative.spec_utils import draft_tp_context logger = logging.getLogger(__name__) @@ -406,8 +409,10 @@ class DraftBlockProposer: ) -> None: if not self._dp_moe_sync or batch.global_num_tokens is None: return - gnt, gnt_logprob = ( - self._draft_block_spec_info.get_spec_adjusted_global_num_tokens(batch) + gnt, gnt_logprob = spec_scale_global_num_tokens( + self._draft_block_spec_info, + batch.global_num_tokens, + batch.global_num_tokens_for_logprob, ) device = self.draft_model_runner.device forward_batch.global_num_tokens_cpu = gnt diff --git a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py index 554a69b2d..2cda3f115 100644 --- a/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_cuda_graph_runner.py @@ -293,12 +293,8 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner): # ----------------------------------------------------------------- def can_run_graph(self, forward_batch: ForwardBatch): if self.require_mlp_tp_gather: - cuda_graph_bs = ( - max(forward_batch.global_num_tokens_cpu) // self.num_tokens_per_req - if self.model_runner.spec_algorithm.is_eagle() - or self.model_runner.spec_algorithm.is_standalone() - else max(forward_batch.global_num_tokens_cpu) - ) + # Raw sync values are per-rank request counts on decode-family rounds. + cuda_graph_bs = max(forward_batch.original_global_num_tokens_cpu) else: cuda_graph_bs = forward_batch.batch_size diff --git a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py index c3bcf9185..1234aece3 100644 --- a/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/eagle_draft_extend_cuda_graph_runner.py @@ -290,12 +290,8 @@ class EAGLEDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): def can_run_graph(self, forward_batch: ForwardBatch): if self.require_mlp_tp_gather: - cuda_graph_bs = ( - max(forward_batch.global_num_tokens_cpu) // self.num_tokens_per_req - if self.model_runner.spec_algorithm.is_eagle() - or self.model_runner.spec_algorithm.is_standalone() - else max(forward_batch.global_num_tokens_cpu) - ) + # Raw sync values are per-rank request counts on decode-family rounds. + cuda_graph_bs = max(forward_batch.original_global_num_tokens_cpu) else: cuda_graph_bs = forward_batch.seq_lens.numel() diff --git a/python/sglang/srt/speculative/eagle_info.py b/python/sglang/srt/speculative/eagle_info.py index dd4d69933..2059ee416 100644 --- a/python/sglang/srt/speculative/eagle_info.py +++ b/python/sglang/srt/speculative/eagle_info.py @@ -1,6 +1,6 @@ import logging from dataclasses import dataclass -from typing import List, Optional, Tuple +from typing import List, Optional import torch @@ -43,13 +43,6 @@ class EagleVerifyInput(SpecInput): self.num_tokens_per_req = self.draft_token_num self.num_tokens_for_logprob_per_req = self.draft_token_num - def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]: - # Keep this override on draft_token_num: eagle_worker_v2.verify() - # re-stamps num_tokens_per_req = num_steps + 1, which diverges from - # the real verify width for topk > 1 trees, and the DP-attention - # global-token scaling must follow the actual tree width. - return self.draft_token_num, self.draft_token_num - @property def max_tree_depth(self) -> int: """Longest root-to-leaf chain of the verify tree, incl. the root; diff --git a/python/sglang/srt/speculative/eagle_worker_v2.py b/python/sglang/srt/speculative/eagle_worker_v2.py index 08713d430..be4800828 100644 --- a/python/sglang/srt/speculative/eagle_worker_v2.py +++ b/python/sglang/srt/speculative/eagle_worker_v2.py @@ -1472,10 +1472,6 @@ class EAGLEWorkerV2(BaseSpecWorker): verify_input: EagleVerifyInput = batch.spec_info record_stream_for_v2_verify(batch, verify_input, fwd_stream) - # Actual-width stamp; equals the real width (draft_token_num) only on - # the topk-1 chain. Tree consumers must read draft_token_num -- see - # EagleVerifyInput.get_spec_adjust_token_coefficient. - verify_input.num_tokens_per_req = self.speculative_num_steps + 1 bs = len(batch.seq_lens) # Batch 1: Target verify diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py index 2ed69aa58..bed3369f5 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_cuda_graph_runner.py @@ -202,8 +202,11 @@ class FrozenKVMTPCudaGraphRunner(DecodeCudaGraphRunner): def can_run_graph(self, forward_batch: ForwardBatch): if self.require_mlp_tp_gather: - cuda_graph_bs = max(forward_batch.global_num_tokens_cpu) // ( - self.topk * self.topk + # Raw sync values are per-rank request counts on decode-family + # rounds; / topk maps the expanded batch to graph-key units + # (mirrors the non-gather branch below). + cuda_graph_bs = ( + max(forward_batch.original_global_num_tokens_cpu) // self.topk ) else: cuda_graph_bs = ( diff --git a/python/sglang/srt/speculative/frozen_kv_mtp_info.py b/python/sglang/srt/speculative/frozen_kv_mtp_info.py index 614907aa8..dc9cc184d 100644 --- a/python/sglang/srt/speculative/frozen_kv_mtp_info.py +++ b/python/sglang/srt/speculative/frozen_kv_mtp_info.py @@ -57,4 +57,6 @@ class FrozenKVMTPVerifyInput(EagleVerifyInput): """Verify input for Frozen-KV MTP.""" def __post_init__(self): - SpecInput.__init__(self, SpecInputType.FROZEN_KV_MTP_VERIFY) + # Run EagleVerifyInput's width auto-fill, then correct the type tag. + super().__post_init__() + self.spec_input_type = SpecInputType.FROZEN_KV_MTP_VERIFY diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index f75439d90..9179902cd 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -201,11 +201,8 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner): def can_run_graph(self, forward_batch: ForwardBatch): if self.require_mlp_tp_gather: - cuda_graph_bs = ( - max(forward_batch.global_num_tokens_cpu) // self.num_tokens_per_req - if self.model_runner.spec_algorithm.is_eagle() - else max(forward_batch.global_num_tokens_cpu) - ) + # Raw sync values are per-rank request counts on decode-family rounds. + cuda_graph_bs = max(forward_batch.original_global_num_tokens_cpu) else: cuda_graph_bs = forward_batch.seq_lens.numel() diff --git a/python/sglang/srt/speculative/spec_info.py b/python/sglang/srt/speculative/spec_info.py index 5fe6f1066..c9d27b760 100644 --- a/python/sglang/srt/speculative/spec_info.py +++ b/python/sglang/srt/speculative/spec_info.py @@ -350,18 +350,22 @@ class SpecInput(ABC): SpecInputType.NGRAM_VERIFY, } - def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]: - return self.num_tokens_per_req, self.num_tokens_for_logprob_per_req - def get_spec_adjusted_global_num_tokens( - self, batch: ScheduleBatch - ) -> Tuple[List[int], List[int]]: - c1, c2 = self.get_spec_adjust_token_coefficient() - global_num_tokens = [x * c1 for x in batch.global_num_tokens] - global_num_tokens_for_logprob = [ - x * c2 for x in batch.global_num_tokens_for_logprob - ] - return global_num_tokens, global_num_tokens_for_logprob +def spec_scale_global_num_tokens( + spec_info: SpecInput, + global_num_tokens: List[int], + global_num_tokens_for_logprob: List[int], +) -> Tuple[List[int], List[int]]: + """Scale the raw per-rank sync values (request counts on decode-family + rounds) into this forward's token units using the spec input's uniform + per-request widths.""" + return ( + [x * spec_info.num_tokens_per_req for x in global_num_tokens], + [ + x * spec_info.num_tokens_for_logprob_per_req + for x in global_num_tokens_for_logprob + ], + ) def create_dummy_verify_input(