[Spec] Centralize dummy verify-input capture; add carries_draft_hidden_states (#28032)
This commit is contained in:
@@ -1106,12 +1106,12 @@ class Scheduler(
|
|||||||
buffer_size,
|
buffer_size,
|
||||||
hidden_size=(
|
hidden_size=(
|
||||||
model_config.spec_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
|
else 16 # minimal padding size for RDMA
|
||||||
),
|
),
|
||||||
hidden_states_dtype=(
|
hidden_states_dtype=(
|
||||||
model_config.dtype
|
model_config.dtype
|
||||||
if self.spec_algorithm.is_eagle()
|
if self.spec_algorithm.carries_draft_hidden_states()
|
||||||
else torch.float32
|
else torch.float32
|
||||||
),
|
),
|
||||||
custom_mem_pool=self.token_to_kv_pool_allocator.get_kvcache().maybe_get_custom_mem_pool(),
|
custom_mem_pool=self.token_to_kv_pool_allocator.get_kvcache().maybe_get_custom_mem_pool(),
|
||||||
@@ -1159,14 +1159,12 @@ class Scheduler(
|
|||||||
buffer_size,
|
buffer_size,
|
||||||
hidden_size=(
|
hidden_size=(
|
||||||
model_config.spec_hidden_size
|
model_config.spec_hidden_size
|
||||||
if self.spec_algorithm.is_eagle()
|
if self.spec_algorithm.carries_draft_hidden_states()
|
||||||
or self.spec_algorithm.is_standalone()
|
|
||||||
else 16 # minimal padding size for RDMA
|
else 16 # minimal padding size for RDMA
|
||||||
),
|
),
|
||||||
hidden_states_dtype=(
|
hidden_states_dtype=(
|
||||||
model_config.dtype
|
model_config.dtype
|
||||||
if self.spec_algorithm.is_eagle()
|
if self.spec_algorithm.carries_draft_hidden_states()
|
||||||
or self.spec_algorithm.is_standalone()
|
|
||||||
else torch.float32
|
else torch.float32
|
||||||
),
|
),
|
||||||
custom_mem_pool=self.token_to_kv_pool_allocator.get_kvcache().maybe_get_custom_mem_pool(),
|
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,
|
get_global_server_args,
|
||||||
set_global_server_args_for_scheduler,
|
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.base import TopkCaptureOutput
|
||||||
from sglang.srt.state_capturer.indexer_topk import (
|
from sglang.srt.state_capturer.indexer_topk import (
|
||||||
create_indexer_capturer,
|
create_indexer_capturer,
|
||||||
@@ -2737,29 +2740,16 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
global_dp_buffer_len = None
|
global_dp_buffer_len = None
|
||||||
global_num_tokens_cpu = None
|
global_num_tokens_cpu = None
|
||||||
|
|
||||||
def get_spec_info():
|
spec_info = create_dummy_verify_input(
|
||||||
spec_info = None
|
self.spec_algorithm,
|
||||||
if self.spec_algorithm.is_eagle() or self.spec_algorithm.is_standalone():
|
self.server_args,
|
||||||
from sglang.srt.speculative.eagle_info import EagleVerifyInput
|
buffers.custom_mask,
|
||||||
|
num_tokens_per_bs,
|
||||||
if self.is_draft_worker:
|
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,
|
|
||||||
)
|
)
|
||||||
|
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
|
# MTP models (e.g. deepseek_nextn) read spec_info.hidden_states
|
||||||
# during forward; provide a dummy so warmup doesn't crash.
|
# during forward; provide a dummy so warmup doesn't crash.
|
||||||
spec_info.hidden_states = torch.zeros(
|
spec_info.hidden_states = torch.zeros(
|
||||||
@@ -2767,39 +2757,6 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
dtype=self.dtype,
|
dtype=self.dtype,
|
||||||
device=self.device,
|
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()
|
|
||||||
if capture_hidden_mode != CaptureHiddenMode.FULL:
|
if capture_hidden_mode != CaptureHiddenMode.FULL:
|
||||||
capture_hidden_mode = (
|
capture_hidden_mode = (
|
||||||
spec_info.capture_hidden_mode if spec_info else CaptureHiddenMode.NULL
|
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."""
|
per-topk page rounding; see get_alloc_len_per_decode."""
|
||||||
return not self.is_ngram()
|
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(
|
def create_future_map(
|
||||||
self,
|
self,
|
||||||
device: torch.device,
|
device: torch.device,
|
||||||
@@ -266,3 +271,66 @@ class SpecInput(ABC):
|
|||||||
x * c2 for x in batch.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
|
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
|
||||||
|
|||||||
Reference in New Issue
Block a user