From caf59759ea043c96a4aa0cc7ac84f2f516b170d8 Mon Sep 17 00:00:00 2001 From: Liangsheng Yin Date: Fri, 12 Jun 2026 14:23:29 -0700 Subject: [PATCH] [Spec] Centralize dummy verify-input capture; add `carries_draft_hidden_states` (#28032) --- python/sglang/srt/managers/scheduler.py | 10 +-- .../sglang/srt/model_executor/model_runner.py | 85 +++++-------------- python/sglang/srt/speculative/spec_info.py | 68 +++++++++++++++ 3 files changed, 93 insertions(+), 70 deletions(-) diff --git a/python/sglang/srt/managers/scheduler.py b/python/sglang/srt/managers/scheduler.py index fa95f870c..24f2ead54 100644 --- a/python/sglang/srt/managers/scheduler.py +++ b/python/sglang/srt/managers/scheduler.py @@ -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(), diff --git a/python/sglang/srt/model_executor/model_runner.py b/python/sglang/srt/model_executor/model_runner.py index 9f43032d8..579f8353b 100644 --- a/python/sglang/srt/model_executor/model_runner.py +++ b/python/sglang/srt/model_executor/model_runner.py @@ -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 diff --git a/python/sglang/srt/speculative/spec_info.py b/python/sglang/srt/speculative/spec_info.py index 92102c825..3fe49c2fb 100644 --- a/python/sglang/srt/speculative/spec_info.py +++ b/python/sglang/srt/speculative/spec_info.py @@ -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