[Spec] Converge DP-attention spec width scaling onto num_tokens_per_req (#31244)

This commit is contained in:
Liangsheng Yin
2026-07-15 17:57:08 -07:00
committed by GitHub
parent 5d004a20c5
commit 1cc9493747
12 changed files with 54 additions and 57 deletions
@@ -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
)
@@ -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
+2 -2
View File
@@ -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(
@@ -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
@@ -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
@@ -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()
+1 -8
View File
@@ -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;
@@ -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
@@ -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 = (
@@ -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
@@ -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()
+15 -11
View File
@@ -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(