[Spec] Extract shared draft worker construction and generalize draft sampler capture (#30857)

This commit is contained in:
Liangsheng Yin
2026-07-11 14:34:09 -05:00
committed by GitHub
parent 348e6fd29b
commit d8ef76682e
8 changed files with 271 additions and 99 deletions
@@ -147,6 +147,9 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None:
def _handle_dflash(server_args: ServerArgs) -> None:
from sglang.srt.arg_groups.overrides import resolved_view
if not server_args.device.startswith("cuda"):
raise ValueError("DFLASH speculative decoding only supports CUDA device.")
if resolved_view(server_args).enable_dp_attention:
raise ValueError(
"Currently DFLASH speculative decoding does not support dp attention."
@@ -390,6 +390,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
self.spec_algorithm = SpeculativeAlgorithm.from_string(
server_args.speculative_algorithm
)
self.capture_tail_hooks = []
self.page_size = server_args.page_size
self.req_to_token_pool = req_to_token_pool
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
@@ -820,21 +820,8 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
forward_batch,
**kwargs,
)
dflash_sampler = getattr(
self.model_runner, "dflash_draft_sampler", None
)
if dflash_sampler is not None:
# Must be captured here, or replay leaves a stale output buffer
# the worker would read as valid tokens -- fail loudly instead.
if (
not isinstance(out, LogitsProcessorOutput)
or out.hidden_states is None
):
raise RuntimeError(
"DFLASH draft sampler set but the draft forward has no "
"hidden_states to capture into the graph."
)
dflash_sampler(out.hidden_states)
for capture_hook in self.model_runner.capture_tail_hooks:
capture_hook(self, out, forward_batch, num_tokens)
return out
self.deepep_adapter.capture(is_extend_in_batch=False)
+6
View File
@@ -2035,6 +2035,9 @@ class DeepseekV4Model(nn.Module):
if self.dsa_enable_prefill_cp:
self.cp_size = get_parallel().attn_cp_size
def get_input_embeddings(self) -> nn.Module:
return self.embed_tokens
def hc_head(
self,
x: torch.Tensor,
@@ -2346,6 +2349,9 @@ class DeepseekV4ForCausalLM(nn.Module):
def routed_experts_weights_of_layer(self):
return self._routed_experts_weights_of_layer.value
def get_input_embeddings(self) -> nn.Module:
return self.model.get_input_embeddings()
def determine_num_fused_shared_experts(self):
self.num_fused_shared_experts = 0
if get_server_args().disable_shared_experts_fusion:
+3 -3
View File
@@ -39,12 +39,12 @@ class DFlashVerifyInput(SpecInput):
capture_hidden_mode: CaptureHiddenMode = CaptureHiddenMode.FULL
# Shape info for padding (e.g., DP attention / CUDA graph).
num_tokens_per_batch: int = -1
num_tokens_per_req: int = -1
def __post_init__(self):
super().__init__(spec_input_type=SpecInputType.DFLASH_VERIFY)
if self.num_tokens_per_batch == -1:
self.num_tokens_per_batch = int(self.draft_token_num)
if self.num_tokens_per_req == -1:
self.num_tokens_per_req = int(self.draft_token_num)
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
return self.draft_token_num, self.draft_token_num
+62 -42
View File
@@ -673,9 +673,69 @@ def compute_dflash_sampling_correct_drafts_and_bonus(
dtype=torch.float32,
)
target_probs = build_dflash_verify_target_probs(
next_token_logits=next_token_logits,
sampling_info=sampling_info,
draft_token_num=draft_token_num,
bs=bs,
max_top_k=max_top_k,
uniform_top_k_value=uniform_top_k_value,
use_sparse_topk=use_sparse_topk,
)
draft_probs = torch.zeros_like(target_probs)
(
retrieve_index,
retrieve_next_token,
retrieve_next_sibling,
predicts,
accept_index,
accept_token_num,
) = _get_or_create_chain_verify_buffers(
bs=bs,
draft_token_num=draft_token_num,
device=device,
)
candidates_i64 = (
candidates if candidates.dtype == torch.int64 else candidates.to(torch.int64)
)
tree_speculative_sampling_target_only(
predicts=predicts,
accept_index=accept_index,
accept_token_num=accept_token_num,
candidates=candidates_i64,
retrive_index=retrieve_index,
retrive_next_token=retrieve_next_token,
retrive_next_sibling=retrieve_next_sibling,
uniform_samples=uniform_samples,
uniform_samples_for_final_sampling=uniform_samples_for_final_sampling,
target_probs=target_probs,
draft_probs=draft_probs,
threshold_single=threshold_single,
threshold_acc=threshold_acc,
deterministic=True,
)
correct_len = accept_token_num
row_ids = torch.arange(bs, dtype=torch.long, device=device)
accept_pos = accept_index[row_ids, correct_len.to(torch.long)].to(torch.long)
bonus = predicts[accept_pos].to(torch.int64)
return correct_len, bonus
def build_dflash_verify_target_probs(
*,
next_token_logits: torch.Tensor,
sampling_info: Any,
draft_token_num: int,
bs: int,
max_top_k: Optional[int] = None,
uniform_top_k_value: Optional[int] = None,
use_sparse_topk: bool = True,
) -> torch.Tensor:
device = next_token_logits.device
need_top_k = bool(getattr(sampling_info, "need_top_k_sampling", True))
need_top_p = bool(getattr(sampling_info, "need_top_p_sampling", False))
# Build target distribution once over all verify rows.
expanded_temperature = torch.repeat_interleave(
sampling_info.temperatures, draft_token_num, dim=0
)
@@ -730,47 +790,7 @@ def compute_dflash_sampling_correct_drafts_and_bonus(
target_probs,
torch.repeat_interleave(sampling_info.top_ps, draft_token_num, dim=0),
)
target_probs = target_probs.view(bs, draft_token_num, -1).contiguous()
draft_probs = torch.zeros_like(target_probs)
(
retrieve_index,
retrieve_next_token,
retrieve_next_sibling,
predicts,
accept_index,
accept_token_num,
) = _get_or_create_chain_verify_buffers(
bs=bs,
draft_token_num=draft_token_num,
device=device,
)
candidates_i64 = (
candidates if candidates.dtype == torch.int64 else candidates.to(torch.int64)
)
tree_speculative_sampling_target_only(
predicts=predicts,
accept_index=accept_index,
accept_token_num=accept_token_num,
candidates=candidates_i64,
# kwarg LHS retained as `retrive_*` to match sgl_kernel op schema.
retrive_index=retrieve_index,
retrive_next_token=retrieve_next_token,
retrive_next_sibling=retrieve_next_sibling,
uniform_samples=uniform_samples,
uniform_samples_for_final_sampling=uniform_samples_for_final_sampling,
target_probs=target_probs,
draft_probs=draft_probs,
threshold_single=threshold_single,
threshold_acc=threshold_acc,
deterministic=True,
)
correct_len = accept_token_num
row_ids = torch.arange(bs, dtype=torch.long, device=device)
accept_pos = accept_index[row_ids, correct_len.to(torch.long)].to(torch.long)
bonus = predicts[accept_pos].to(torch.int64)
return correct_len, bonus
return target_probs.view(bs, draft_token_num, -1).contiguous()
def validate_dflash_request(req: Req, enable_overlap: bool) -> Optional[str]:
@@ -32,6 +32,13 @@ from sglang.srt.speculative.dflash_utils import (
is_dflash_sampling_verify_available,
parse_dflash_draft_config,
)
from sglang.srt.speculative.draft_worker_common import (
build_block_pos_offsets,
build_draft_tp_worker,
make_draft_block_spec_info,
make_draft_input_v2,
make_draft_sampler_capture_hook,
)
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_utils import assign_req_to_token_pool_func
from sglang.srt.utils import get_available_gpu_memory, is_cuda, is_hip, is_npu
@@ -74,7 +81,7 @@ class _DflashDraftSampler:
device=weight.device,
)
def __call__(self, hidden_states):
def __call__(self, hidden_states, input_ids=None):
# draft tokens are block positions 1: (pos 0 is the seeded bonus token)
bs = hidden_states.shape[0] // self.block_size
hs = hidden_states.view(bs, self.block_size, -1)[:, 1:, :].reshape(
@@ -130,25 +137,22 @@ class DFlashWorkerV2(BaseSpecWorker):
self._warned_sampling_fallback = False
self._logged_first_verify = False
# Draft runner (separate KV cache + attention backend).
self._draft_worker = TpModelWorker(
bundle = build_draft_tp_worker(
server_args=server_args,
gpu_id=gpu_id,
tp_rank=tp_rank,
dp_rank=dp_rank,
moe_ep_rank=moe_ep_rank,
pp_rank=0,
attn_cp_rank=attn_cp_rank,
moe_dp_rank=moe_dp_rank,
dp_rank=dp_rank,
nccl_port=nccl_port,
is_draft_worker=True,
context_length=target_worker.model_runner.model_config.context_len,
target_model_config=target_worker.model_runner.model_config,
algo_label="DFLASH",
)
self.draft_model_runner = self._draft_worker.model_runner
self._draft_worker = bundle.draft_worker
self.draft_model_runner = bundle.draft_model_runner
self._draft_sampler = None
# Keep the same alias that other spec-v2 workers expose.
self._draft_worker.draft_runner = self.draft_model_runner
self.draft_model = self.draft_model_runner.model
self.draft_model = bundle.draft_model
draft_config = parse_dflash_draft_config(
draft_hf_config=self.draft_model_runner.model_config.hf_config
)
@@ -179,7 +183,7 @@ class DFlashWorkerV2(BaseSpecWorker):
if self.tp_rank == 0:
logger.info(
"Initialized DFLASH draft runner. attention_backend=%s, model=%s, block_size=%s, draft_window_size=%s, compact_cache=%s",
server_args.speculative_draft_attention_backend,
bundle.resolved_attention_backend,
self.draft_model.__class__.__name__,
self.block_size,
self.draft_window_size,
@@ -192,8 +196,8 @@ class DFlashWorkerV2(BaseSpecWorker):
self._mask_token_id_override,
)
self._block_pos_offsets = torch.arange(
self.block_size, device=self.device, dtype=torch.int64
self._block_pos_offsets = build_block_pos_offsets(
length=self.block_size, device=self.device
)
self._draft_block_ids_buf: Optional[torch.Tensor] = None # [cap_bs, block_size]
self._draft_block_positions_buf: Optional[torch.Tensor] = (
@@ -207,12 +211,8 @@ class DFlashWorkerV2(BaseSpecWorker):
)
self._draft_block_end_buf: Optional[torch.Tensor] = None # [cap_bs]
self._draft_seq_lens_cpu_buf: Optional[torch.Tensor] = None # [cap_bs] on CPU
self._draft_block_spec_info = DFlashVerifyInput(
draft_token=torch.empty((0,), dtype=torch.long, device=self.device),
positions=torch.empty((0,), dtype=torch.int64, device=self.device),
draft_token_num=int(self.block_size),
custom_mask=None,
capture_hidden_mode=CaptureHiddenMode.NULL,
self._draft_block_spec_info = make_draft_block_spec_info(
draft_token_num=int(self.block_size), device=self.device
)
self._draft_greedy_gathered_max_buf: Optional[torch.Tensor] = None
self._draft_greedy_gathered_ids_buf: Optional[torch.Tensor] = None
@@ -305,7 +305,10 @@ class DFlashWorkerV2(BaseSpecWorker):
if capture_decode_cuda_graph:
# Must run before capture so the draft graph folds the head in.
self._draft_sampler = self._maybe_build_draft_sampler()
self.draft_model_runner.dflash_draft_sampler = self._draft_sampler
if self._draft_sampler is not None:
self.draft_model_runner.capture_tail_hooks.append(
make_draft_sampler_capture_hook(self._draft_sampler)
)
self._draft_worker.init_cuda_graphs(
capture_decode_cuda_graph=capture_decode_cuda_graph
)
@@ -1214,15 +1217,7 @@ class DFlashWorkerV2(BaseSpecWorker):
bonus_tokens: torch.Tensor,
seq_lens: torch.Tensor,
) -> DFlashDraftInputV2:
bs = int(seq_lens.numel())
device = bonus_tokens.device
return DFlashDraftInputV2(
topk_p=torch.empty((bs, 0), device=device, dtype=torch.float32),
topk_index=torch.empty((bs, 0), device=device, dtype=torch.int64),
bonus_tokens=bonus_tokens.to(dtype=torch.int64),
new_seq_lens=seq_lens.to(dtype=torch.int64),
hidden_states=torch.empty((bs, 0), device=device, dtype=torch.float16),
)
return make_draft_input_v2(bonus_tokens=bonus_tokens, new_seq_lens=seq_lens)
def _make_next_draft_input_decode(
self,
@@ -1230,15 +1225,7 @@ class DFlashWorkerV2(BaseSpecWorker):
bonus_tokens: torch.Tensor,
new_seq_lens: torch.Tensor,
) -> DFlashDraftInputV2:
bs = int(new_seq_lens.numel())
device = bonus_tokens.device
return DFlashDraftInputV2(
topk_p=torch.empty((bs, 0), device=device, dtype=torch.float32),
topk_index=torch.empty((bs, 0), device=device, dtype=torch.int64),
bonus_tokens=bonus_tokens.to(dtype=torch.int64),
new_seq_lens=new_seq_lens.to(dtype=torch.int64),
hidden_states=torch.empty((bs, 0), device=device, dtype=torch.float16),
)
return make_draft_input_v2(bonus_tokens=bonus_tokens, new_seq_lens=new_seq_lens)
def forward_batch_generation(
self,
@@ -0,0 +1,168 @@
from __future__ import annotations
import logging
from copy import deepcopy
from typing import TYPE_CHECKING, Optional
import msgspec
import torch
from sglang.srt.layers.logits_processor import LogitsProcessorOutput
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
from sglang.srt.runtime_context import get_context, get_server_args
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.dflash_info import DFlashVerifyInput
from sglang.srt.speculative.dflash_info_v2 import DFlashDraftInputV2
if TYPE_CHECKING:
from sglang.srt.configs.model_config import ModelConfig
from sglang.srt.model_executor.model_runner import ModelRunner
logger = logging.getLogger(__name__)
_SUPPORTED_DRAFT_BACKENDS = ("flashinfer", "fa3", "fa4", "triton", "ascend")
class DraftWorkerBundle(msgspec.Struct, frozen=True):
draft_worker: TpModelWorker
draft_model_runner: ModelRunner
draft_model: torch.nn.Module
resolved_attention_backend: str
def _resolve_draft_attention_backend_fallback(
*, draft_server_args: ServerArgs, algo_label: str
) -> str:
draft_backend = draft_server_args.speculative_draft_attention_backend
if draft_backend is None:
draft_backend, _ = draft_server_args.get_attention_backends()
if draft_backend is None:
return "triton" if torch.version.hip else "flashinfer"
if draft_backend not in _SUPPORTED_DRAFT_BACKENDS:
fallback = "triton" if torch.version.hip else "flashinfer"
logger.warning(
"%s draft worker only supports attention_backend in %s for now, "
"but got %r. Falling back to '%s'.",
algo_label,
_SUPPORTED_DRAFT_BACKENDS,
draft_backend,
fallback,
)
return fallback
return draft_backend
def build_draft_tp_worker(
*,
server_args: ServerArgs,
gpu_id: int,
tp_rank: int,
dp_rank: Optional[int],
moe_ep_rank: int,
attn_cp_rank: int,
moe_dp_rank: int,
nccl_port: int,
target_model_config: ModelConfig,
algo_label: str,
attention_backend_override: Optional[str] = None,
) -> DraftWorkerBundle:
draft_server_args = deepcopy(server_args)
# An override names a draft-specific backend the caller has already
# validated (e.g. a self-drafting architecture); it skips the generic
# supported-backend fallback below.
draft_backend = attention_backend_override or (
_resolve_draft_attention_backend_fallback(
draft_server_args=draft_server_args, algo_label=algo_label
)
)
# Post-resolution ServerArgs rejects bare assignment; route the draft-copy
# adjustments through the audited mutation point. Keep the resolved value
# on speculative_draft_attention_backend: downstream draft-worker logic
# keys on that field (backend selection in _get_attention_backend and the
# fa4-draft KV dtype override in configure_kv_cache_dtype), so nulling it
# would silently skip those paths. context_length keeps the draft aligned
# with the target.
draft_server_args.override(
"draft_worker.build",
skip_tokenizer_init=True,
speculative_draft_attention_backend=draft_backend,
prefill_attention_backend=None,
decode_attention_backend=None,
attention_backend=draft_backend,
context_length=target_model_config.context_len,
)
saved_server_args = get_server_args()
try:
draft_worker = TpModelWorker(
server_args=draft_server_args,
gpu_id=gpu_id,
tp_rank=tp_rank,
moe_ep_rank=moe_ep_rank,
pp_rank=0,
attn_cp_rank=attn_cp_rank,
moe_dp_rank=moe_dp_rank,
dp_rank=dp_rank,
nccl_port=nccl_port,
is_draft_worker=True,
)
finally:
get_context().set_server_args(saved_server_args)
draft_model_runner = draft_worker.model_runner
draft_worker.draft_runner = draft_model_runner
return DraftWorkerBundle(
draft_worker=draft_worker,
draft_model_runner=draft_model_runner,
draft_model=draft_model_runner.model,
resolved_attention_backend=draft_backend,
)
def make_draft_input_v2(
*,
bonus_tokens: torch.Tensor,
new_seq_lens: torch.Tensor,
) -> DFlashDraftInputV2:
bs = int(new_seq_lens.numel())
device = bonus_tokens.device
return DFlashDraftInputV2(
topk_p=torch.empty((bs, 0), device=device, dtype=torch.float32),
topk_index=torch.empty((bs, 0), device=device, dtype=torch.int64),
bonus_tokens=bonus_tokens.to(dtype=torch.int64),
new_seq_lens=new_seq_lens.to(dtype=torch.int64),
hidden_states=torch.empty((bs, 0), device=device, dtype=torch.float16),
)
def make_draft_block_spec_info(
*,
draft_token_num: int,
device: torch.device,
) -> DFlashVerifyInput:
return DFlashVerifyInput(
draft_token=torch.empty((0,), dtype=torch.long, device=device),
positions=torch.empty((0,), dtype=torch.int64, device=device),
draft_token_num=int(draft_token_num),
custom_mask=None,
capture_hidden_mode=CaptureHiddenMode.NULL,
)
def make_draft_sampler_capture_hook(draft_sampler):
def capture_hook(runner, out, forward_batch, num_tokens):
del runner, num_tokens
if not isinstance(out, LogitsProcessorOutput) or out.hidden_states is None:
raise RuntimeError(
"draft sampler set but the draft forward has no "
"hidden_states to capture into the graph."
)
draft_sampler(out.hidden_states, forward_batch.input_ids)
return capture_hook
def build_block_pos_offsets(*, length: int, device: torch.device) -> torch.Tensor:
return torch.arange(int(length), device=device, dtype=torch.int64)