[Spec] Extract shared draft worker construction and generalize draft sampler capture (#30857)
This commit is contained in:
@@ -147,6 +147,9 @@ def handle_speculative_decoding(server_args: ServerArgs) -> None:
|
|||||||
def _handle_dflash(server_args: ServerArgs) -> None:
|
def _handle_dflash(server_args: ServerArgs) -> None:
|
||||||
from sglang.srt.arg_groups.overrides import resolved_view
|
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:
|
if resolved_view(server_args).enable_dp_attention:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
"Currently DFLASH speculative decoding does not support dp attention."
|
"Currently DFLASH speculative decoding does not support dp attention."
|
||||||
|
|||||||
@@ -390,6 +390,7 @@ class ModelRunner(ModelRunnerKVCacheMixin):
|
|||||||
self.spec_algorithm = SpeculativeAlgorithm.from_string(
|
self.spec_algorithm = SpeculativeAlgorithm.from_string(
|
||||||
server_args.speculative_algorithm
|
server_args.speculative_algorithm
|
||||||
)
|
)
|
||||||
|
self.capture_tail_hooks = []
|
||||||
self.page_size = server_args.page_size
|
self.page_size = server_args.page_size
|
||||||
self.req_to_token_pool = req_to_token_pool
|
self.req_to_token_pool = req_to_token_pool
|
||||||
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
|
self.token_to_kv_pool_allocator = token_to_kv_pool_allocator
|
||||||
|
|||||||
@@ -820,21 +820,8 @@ class DecodeCudaGraphRunner(BaseCudaGraphRunner):
|
|||||||
forward_batch,
|
forward_batch,
|
||||||
**kwargs,
|
**kwargs,
|
||||||
)
|
)
|
||||||
dflash_sampler = getattr(
|
for capture_hook in self.model_runner.capture_tail_hooks:
|
||||||
self.model_runner, "dflash_draft_sampler", None
|
capture_hook(self, out, forward_batch, num_tokens)
|
||||||
)
|
|
||||||
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)
|
|
||||||
return out
|
return out
|
||||||
|
|
||||||
self.deepep_adapter.capture(is_extend_in_batch=False)
|
self.deepep_adapter.capture(is_extend_in_batch=False)
|
||||||
|
|||||||
@@ -2035,6 +2035,9 @@ class DeepseekV4Model(nn.Module):
|
|||||||
if self.dsa_enable_prefill_cp:
|
if self.dsa_enable_prefill_cp:
|
||||||
self.cp_size = get_parallel().attn_cp_size
|
self.cp_size = get_parallel().attn_cp_size
|
||||||
|
|
||||||
|
def get_input_embeddings(self) -> nn.Module:
|
||||||
|
return self.embed_tokens
|
||||||
|
|
||||||
def hc_head(
|
def hc_head(
|
||||||
self,
|
self,
|
||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
@@ -2346,6 +2349,9 @@ class DeepseekV4ForCausalLM(nn.Module):
|
|||||||
def routed_experts_weights_of_layer(self):
|
def routed_experts_weights_of_layer(self):
|
||||||
return self._routed_experts_weights_of_layer.value
|
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):
|
def determine_num_fused_shared_experts(self):
|
||||||
self.num_fused_shared_experts = 0
|
self.num_fused_shared_experts = 0
|
||||||
if get_server_args().disable_shared_experts_fusion:
|
if get_server_args().disable_shared_experts_fusion:
|
||||||
|
|||||||
@@ -39,12 +39,12 @@ class DFlashVerifyInput(SpecInput):
|
|||||||
capture_hidden_mode: CaptureHiddenMode = CaptureHiddenMode.FULL
|
capture_hidden_mode: CaptureHiddenMode = CaptureHiddenMode.FULL
|
||||||
|
|
||||||
# Shape info for padding (e.g., DP attention / CUDA graph).
|
# 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):
|
def __post_init__(self):
|
||||||
super().__init__(spec_input_type=SpecInputType.DFLASH_VERIFY)
|
super().__init__(spec_input_type=SpecInputType.DFLASH_VERIFY)
|
||||||
if self.num_tokens_per_batch == -1:
|
if self.num_tokens_per_req == -1:
|
||||||
self.num_tokens_per_batch = int(self.draft_token_num)
|
self.num_tokens_per_req = int(self.draft_token_num)
|
||||||
|
|
||||||
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
|
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
|
||||||
return self.draft_token_num, self.draft_token_num
|
return self.draft_token_num, self.draft_token_num
|
||||||
|
|||||||
@@ -673,9 +673,69 @@ def compute_dflash_sampling_correct_drafts_and_bonus(
|
|||||||
dtype=torch.float32,
|
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_k = bool(getattr(sampling_info, "need_top_k_sampling", True))
|
||||||
need_top_p = bool(getattr(sampling_info, "need_top_p_sampling", False))
|
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(
|
expanded_temperature = torch.repeat_interleave(
|
||||||
sampling_info.temperatures, draft_token_num, dim=0
|
sampling_info.temperatures, draft_token_num, dim=0
|
||||||
)
|
)
|
||||||
@@ -730,47 +790,7 @@ def compute_dflash_sampling_correct_drafts_and_bonus(
|
|||||||
target_probs,
|
target_probs,
|
||||||
torch.repeat_interleave(sampling_info.top_ps, draft_token_num, dim=0),
|
torch.repeat_interleave(sampling_info.top_ps, draft_token_num, dim=0),
|
||||||
)
|
)
|
||||||
target_probs = target_probs.view(bs, draft_token_num, -1).contiguous()
|
return 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
|
|
||||||
|
|
||||||
|
|
||||||
def validate_dflash_request(req: Req, enable_overlap: bool) -> Optional[str]:
|
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,
|
is_dflash_sampling_verify_available,
|
||||||
parse_dflash_draft_config,
|
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_info import SpeculativeAlgorithm
|
||||||
from sglang.srt.speculative.spec_utils import assign_req_to_token_pool_func
|
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
|
from sglang.srt.utils import get_available_gpu_memory, is_cuda, is_hip, is_npu
|
||||||
@@ -74,7 +81,7 @@ class _DflashDraftSampler:
|
|||||||
device=weight.device,
|
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)
|
# draft tokens are block positions 1: (pos 0 is the seeded bonus token)
|
||||||
bs = hidden_states.shape[0] // self.block_size
|
bs = hidden_states.shape[0] // self.block_size
|
||||||
hs = hidden_states.view(bs, self.block_size, -1)[:, 1:, :].reshape(
|
hs = hidden_states.view(bs, self.block_size, -1)[:, 1:, :].reshape(
|
||||||
@@ -130,25 +137,22 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
self._warned_sampling_fallback = False
|
self._warned_sampling_fallback = False
|
||||||
self._logged_first_verify = False
|
self._logged_first_verify = False
|
||||||
|
|
||||||
# Draft runner (separate KV cache + attention backend).
|
bundle = build_draft_tp_worker(
|
||||||
self._draft_worker = TpModelWorker(
|
|
||||||
server_args=server_args,
|
server_args=server_args,
|
||||||
gpu_id=gpu_id,
|
gpu_id=gpu_id,
|
||||||
tp_rank=tp_rank,
|
tp_rank=tp_rank,
|
||||||
|
dp_rank=dp_rank,
|
||||||
moe_ep_rank=moe_ep_rank,
|
moe_ep_rank=moe_ep_rank,
|
||||||
pp_rank=0,
|
|
||||||
attn_cp_rank=attn_cp_rank,
|
attn_cp_rank=attn_cp_rank,
|
||||||
moe_dp_rank=moe_dp_rank,
|
moe_dp_rank=moe_dp_rank,
|
||||||
dp_rank=dp_rank,
|
|
||||||
nccl_port=nccl_port,
|
nccl_port=nccl_port,
|
||||||
is_draft_worker=True,
|
target_model_config=target_worker.model_runner.model_config,
|
||||||
context_length=target_worker.model_runner.model_config.context_len,
|
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
|
self._draft_sampler = None
|
||||||
# Keep the same alias that other spec-v2 workers expose.
|
self.draft_model = bundle.draft_model
|
||||||
self._draft_worker.draft_runner = self.draft_model_runner
|
|
||||||
self.draft_model = self.draft_model_runner.model
|
|
||||||
draft_config = parse_dflash_draft_config(
|
draft_config = parse_dflash_draft_config(
|
||||||
draft_hf_config=self.draft_model_runner.model_config.hf_config
|
draft_hf_config=self.draft_model_runner.model_config.hf_config
|
||||||
)
|
)
|
||||||
@@ -179,7 +183,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
if self.tp_rank == 0:
|
if self.tp_rank == 0:
|
||||||
logger.info(
|
logger.info(
|
||||||
"Initialized DFLASH draft runner. attention_backend=%s, model=%s, block_size=%s, draft_window_size=%s, compact_cache=%s",
|
"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.draft_model.__class__.__name__,
|
||||||
self.block_size,
|
self.block_size,
|
||||||
self.draft_window_size,
|
self.draft_window_size,
|
||||||
@@ -192,8 +196,8 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
self._mask_token_id_override,
|
self._mask_token_id_override,
|
||||||
)
|
)
|
||||||
|
|
||||||
self._block_pos_offsets = torch.arange(
|
self._block_pos_offsets = build_block_pos_offsets(
|
||||||
self.block_size, device=self.device, dtype=torch.int64
|
length=self.block_size, device=self.device
|
||||||
)
|
)
|
||||||
self._draft_block_ids_buf: Optional[torch.Tensor] = None # [cap_bs, block_size]
|
self._draft_block_ids_buf: Optional[torch.Tensor] = None # [cap_bs, block_size]
|
||||||
self._draft_block_positions_buf: Optional[torch.Tensor] = (
|
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_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_seq_lens_cpu_buf: Optional[torch.Tensor] = None # [cap_bs] on CPU
|
||||||
self._draft_block_spec_info = DFlashVerifyInput(
|
self._draft_block_spec_info = make_draft_block_spec_info(
|
||||||
draft_token=torch.empty((0,), dtype=torch.long, device=self.device),
|
draft_token_num=int(self.block_size), 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_greedy_gathered_max_buf: Optional[torch.Tensor] = None
|
self._draft_greedy_gathered_max_buf: Optional[torch.Tensor] = None
|
||||||
self._draft_greedy_gathered_ids_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:
|
if capture_decode_cuda_graph:
|
||||||
# Must run before capture so the draft graph folds the head in.
|
# Must run before capture so the draft graph folds the head in.
|
||||||
self._draft_sampler = self._maybe_build_draft_sampler()
|
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(
|
self._draft_worker.init_cuda_graphs(
|
||||||
capture_decode_cuda_graph=capture_decode_cuda_graph
|
capture_decode_cuda_graph=capture_decode_cuda_graph
|
||||||
)
|
)
|
||||||
@@ -1214,15 +1217,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
bonus_tokens: torch.Tensor,
|
bonus_tokens: torch.Tensor,
|
||||||
seq_lens: torch.Tensor,
|
seq_lens: torch.Tensor,
|
||||||
) -> DFlashDraftInputV2:
|
) -> DFlashDraftInputV2:
|
||||||
bs = int(seq_lens.numel())
|
return make_draft_input_v2(bonus_tokens=bonus_tokens, new_seq_lens=seq_lens)
|
||||||
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),
|
|
||||||
)
|
|
||||||
|
|
||||||
def _make_next_draft_input_decode(
|
def _make_next_draft_input_decode(
|
||||||
self,
|
self,
|
||||||
@@ -1230,15 +1225,7 @@ class DFlashWorkerV2(BaseSpecWorker):
|
|||||||
bonus_tokens: torch.Tensor,
|
bonus_tokens: torch.Tensor,
|
||||||
new_seq_lens: torch.Tensor,
|
new_seq_lens: torch.Tensor,
|
||||||
) -> DFlashDraftInputV2:
|
) -> DFlashDraftInputV2:
|
||||||
bs = int(new_seq_lens.numel())
|
return make_draft_input_v2(bonus_tokens=bonus_tokens, new_seq_lens=new_seq_lens)
|
||||||
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 forward_batch_generation(
|
def forward_batch_generation(
|
||||||
self,
|
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)
|
||||||
Reference in New Issue
Block a user