[Spec] Centralize dummy verify-input capture; add carries_draft_hidden_states (#28032)

This commit is contained in:
Liangsheng Yin
2026-06-12 14:23:29 -07:00
committed by GitHub
parent 65d76bd3f6
commit caf59759ea
3 changed files with 93 additions and 70 deletions
+4 -6
View File
@@ -1106,12 +1106,12 @@ class Scheduler(
buffer_size,
hidden_size=(
model_config.spec_hidden_size
if self.spec_algorithm.is_eagle()
if self.spec_algorithm.carries_draft_hidden_states()
else 16 # minimal padding size for RDMA
),
hidden_states_dtype=(
model_config.dtype
if self.spec_algorithm.is_eagle()
if self.spec_algorithm.carries_draft_hidden_states()
else torch.float32
),
custom_mem_pool=self.token_to_kv_pool_allocator.get_kvcache().maybe_get_custom_mem_pool(),
@@ -1159,14 +1159,12 @@ class Scheduler(
buffer_size,
hidden_size=(
model_config.spec_hidden_size
if self.spec_algorithm.is_eagle()
or self.spec_algorithm.is_standalone()
if self.spec_algorithm.carries_draft_hidden_states()
else 16 # minimal padding size for RDMA
),
hidden_states_dtype=(
model_config.dtype
if self.spec_algorithm.is_eagle()
or self.spec_algorithm.is_standalone()
if self.spec_algorithm.carries_draft_hidden_states()
else torch.float32
),
custom_mem_pool=self.token_to_kv_pool_allocator.get_kvcache().maybe_get_custom_mem_pool(),
@@ -191,7 +191,10 @@ from sglang.srt.server_args import (
get_global_server_args,
set_global_server_args_for_scheduler,
)
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_info import (
SpeculativeAlgorithm,
create_dummy_verify_input,
)
from sglang.srt.state_capturer.base import TopkCaptureOutput
from sglang.srt.state_capturer.indexer_topk import (
create_indexer_capturer,
@@ -2737,69 +2740,23 @@ class ModelRunner(ModelRunnerKVCacheMixin):
global_dp_buffer_len = None
global_num_tokens_cpu = None
def get_spec_info():
spec_info = None
if self.spec_algorithm.is_eagle() or self.spec_algorithm.is_standalone():
from sglang.srt.speculative.eagle_info import EagleVerifyInput
if self.is_draft_worker:
raise RuntimeError("This should not happen.")
else:
spec_info = EagleVerifyInput(
draft_token=None,
custom_mask=buffers.custom_mask,
positions=None,
retrieve_index=None,
retrieve_next_token=None,
retrieve_next_sibling=None,
retrieve_cum_len=None,
spec_steps=self.server_args.speculative_num_steps,
topk=self.server_args.speculative_eagle_topk,
draft_token_num=self.server_args.speculative_num_draft_tokens,
capture_hidden_mode=CaptureHiddenMode.FULL,
seq_lens_sum=None,
seq_lens_cpu=None,
)
# MTP models (e.g. deepseek_nextn) read spec_info.hidden_states
# during forward; provide a dummy so warmup doesn't crash.
spec_info.hidden_states = torch.zeros(
(num_tokens, self.model_config.hidden_size),
dtype=self.dtype,
device=self.device,
)
elif self.spec_algorithm.is_dflash():
from sglang.srt.speculative.dflash_info import DFlashVerifyInput
# Dummy warmup only needs shape metadata; avoid forcing custom-mask mode.
spec_info = DFlashVerifyInput(
draft_token=None,
positions=None,
draft_token_num=self.server_args.speculative_num_draft_tokens,
custom_mask=None,
capture_hidden_mode=(
CaptureHiddenMode.NULL
if self.is_draft_worker
else CaptureHiddenMode.FULL
),
)
elif self.spec_algorithm.is_ngram():
from sglang.srt.speculative.ngram_info import NgramVerifyInput
spec_info = NgramVerifyInput(
draft_token=None,
custom_mask=buffers.custom_mask,
positions=None,
retrieve_index=None,
retrieve_next_token=None,
retrieve_next_sibling=None,
draft_token_num=num_tokens_per_bs,
)
spec_info.capture_hidden_mode = CaptureHiddenMode.NULL
return spec_info
spec_info = get_spec_info()
spec_info = create_dummy_verify_input(
self.spec_algorithm,
self.server_args,
buffers.custom_mask,
num_tokens_per_bs,
self.is_draft_worker,
)
if spec_info is not None and (
self.spec_algorithm.is_eagle() or self.spec_algorithm.is_standalone()
):
# MTP models (e.g. deepseek_nextn) read spec_info.hidden_states
# during forward; provide a dummy so warmup doesn't crash.
spec_info.hidden_states = torch.zeros(
(num_tokens, self.model_config.hidden_size),
dtype=self.dtype,
device=self.device,
)
if capture_hidden_mode != CaptureHiddenMode.FULL:
capture_hidden_mode = (
spec_info.capture_hidden_mode if spec_info else CaptureHiddenMode.NULL
@@ -124,6 +124,11 @@ class SpeculativeAlgorithm(Enum):
per-topk page rounding; see get_alloc_len_per_decode."""
return not self.is_ngram()
def carries_draft_hidden_states(self) -> bool:
"""Whether the disagg prefill->decode transfer carries draft hidden
states (EAGLE-family only; STANDALONE's vanilla draft ignores them)."""
return self.is_eagle()
def create_future_map(
self,
device: torch.device,
@@ -266,3 +271,66 @@ class SpecInput(ABC):
x * c2 for x in batch.global_num_tokens_for_logprob
]
return global_num_tokens, global_num_tokens_for_logprob
def create_dummy_verify_input(
spec_algorithm: SpeculativeAlgorithm,
server_args: ServerArgs,
custom_mask: torch.Tensor,
num_tokens_per_bs: int,
is_draft_worker: bool,
) -> Optional[SpecInput]:
"""Dummy verify ``SpecInput`` for CUDA-graph capture (per-algorithm dispatch)."""
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
spec_info = None
if spec_algorithm.is_eagle() or spec_algorithm.is_standalone():
from sglang.srt.speculative.eagle_info import EagleVerifyInput
if is_draft_worker:
raise RuntimeError("This should not happen.")
else:
spec_info = EagleVerifyInput(
draft_token=None,
custom_mask=custom_mask,
positions=None,
retrieve_index=None,
retrieve_next_token=None,
retrieve_next_sibling=None,
retrieve_cum_len=None,
spec_steps=server_args.speculative_num_steps,
topk=server_args.speculative_eagle_topk,
draft_token_num=server_args.speculative_num_draft_tokens,
capture_hidden_mode=CaptureHiddenMode.FULL,
seq_lens_sum=None,
seq_lens_cpu=None,
)
elif spec_algorithm.is_dflash():
from sglang.srt.speculative.dflash_info import DFlashVerifyInput
# Dummy warmup only needs shape metadata; avoid forcing custom-mask mode.
spec_info = DFlashVerifyInput(
draft_token=None,
positions=None,
draft_token_num=server_args.speculative_num_draft_tokens,
custom_mask=None,
capture_hidden_mode=(
CaptureHiddenMode.NULL if is_draft_worker else CaptureHiddenMode.FULL
),
)
elif spec_algorithm.is_ngram():
from sglang.srt.speculative.ngram_info import NgramVerifyInput
spec_info = NgramVerifyInput(
draft_token=None,
custom_mask=custom_mask,
positions=None,
retrieve_index=None,
retrieve_next_token=None,
retrieve_next_sibling=None,
draft_token_num=num_tokens_per_bs,
)
spec_info.capture_hidden_mode = CaptureHiddenMode.NULL
return spec_info