[mtp] add rejection sampling for speculative decoding (#26312)

Co-authored-by: lyc508653 <lyc508653@alibaba-inc.com>
Co-authored-by: Qiaolin-Yu <liin1211@outlook.com>
Co-authored-by: Huiqiang Jiang <30883354+iofu728@users.noreply.github.com>
Co-authored-by: Yi Zhang <25844240+yizhang2077@users.noreply.github.com>
Co-authored-by: Yizhong Cao <114661107+cao1zhg@users.noreply.github.com>
This commit is contained in:
Yucheng Li
2026-06-20 15:10:42 -07:00
committed by GitHub
co-authored by lyc508653 Qiaolin-Yu Huiqiang Jiang Yi Zhang Yizhong Cao
parent 95fb1ef697
commit f42ec350b4
11 changed files with 566 additions and 20 deletions
@@ -338,6 +338,47 @@ def _handle_eagle_family(server_args: ServerArgs) -> None:
"trtllm_mha backend only supports topk = 1 for speculative decoding."
)
if server_args.speculative_use_rejection_sampling:
# Resolved alias by now: NEXTN -> EAGLE, Gemma4 draft -> FROZEN_KV_MTP.
# Only the EAGLE/EAGLE3 draft workers emit a target-vocab proposal that
# the rejection-sampling kernel consumes; everything else (STANDALONE,
# FROZEN_KV_MTP, NGRAM, DFLASH) is unsupported.
if server_args.speculative_algorithm not in ("EAGLE", "EAGLE3"):
raise NotImplementedError(
"--speculative-use-rejection-sampling is only supported for "
"EAGLE / EAGLE3 / NEXTN, not "
f"speculative_algorithm={server_args.speculative_algorithm}."
)
if server_args.speculative_eagle_topk != 1:
raise ValueError(
"--speculative-use-rejection-sampling requires --speculative-eagle-topk=1."
)
if (
server_args.speculative_accept_threshold_single != 1.0
or server_args.speculative_accept_threshold_acc != 1.0
):
raise ValueError(
"--speculative-use-rejection-sampling is incompatible with "
"--speculative-accept-threshold-single / "
"--speculative-accept-threshold-acc; rejection sampling ignores "
"the accept thresholds."
)
if server_args.enable_deterministic_inference:
raise ValueError(
"--speculative-use-rejection-sampling is incompatible with "
"--enable-deterministic-inference; the sampling kernel draws "
"coins from the global RNG and is not batch-invariant."
)
if server_args.enable_multi_layer_eagle:
raise NotImplementedError(
"--speculative-use-rejection-sampling is not supported with "
"multi-layer EAGLE (--enable-multi-layer-eagle)."
)
logger.info(
"Rejection sampling is enabled for speculative decoding "
"(speculative_use_rejection_sampling=True)."
)
if (
server_args.speculative_eagle_topk == 1
and server_args.speculative_num_draft_tokens
@@ -192,6 +192,15 @@ class FutureMap:
device=self.device,
)
self.draft_probs_buf = None
if getattr(draft_input, "draft_probs", None) is not None:
draft_probs0 = draft_input.draft_probs[0]
self.draft_probs_buf = torch.empty(
(self.req_pool_size, *draft_probs0.shape),
dtype=draft_probs0.dtype,
device=self.device,
)
def _resolve_spec_extras(self, batch: ScheduleBatch) -> None:
if self.spec_algo.is_ngram():
# FIXME: remove once precomputed draft is supported.
@@ -232,6 +241,8 @@ class FutureMap:
draft_input.bonus_tokens = bonus_tokens
if hidden_states is not None:
draft_input.hidden_states = hidden_states
if self.draft_probs_buf is not None and draft_input.draft_probs is not None:
draft_input.draft_probs = self.draft_probs_buf[indices]
elif self.need_bonus_tokens:
draft_input.bonus_tokens = self.output_tokens_buf[indices]
if self.need_hidden_states and not self.need_topk:
@@ -352,3 +363,5 @@ class FutureMap:
self.hidden_states_buf[indices] = draft_input.hidden_states.to(
self.hidden_states_buf.dtype
)
if self.draft_probs_buf is not None and draft_input.draft_probs is not None:
self.draft_probs_buf[indices] = draft_input.draft_probs
@@ -1336,6 +1336,10 @@ class ForwardBatch(ForwardBatchDeepSeekMHAMixin):
spec_info.topk_index = self._pad_tensor_to_size(
spec_info.topk_index, bs
)
if getattr(spec_info, "draft_probs", None) is not None:
spec_info.draft_probs = self._pad_tensor_to_size(
spec_info.draft_probs, bs
)
if getattr(spec_info, "num_correct_drafts", None) is not None:
spec_info.num_correct_drafts = self._pad_tensor_to_size(
spec_info.num_correct_drafts, bs
+7
View File
@@ -620,6 +620,7 @@ class ServerArgs:
speculative_dflash_block_size: Optional[int] = None
speculative_accept_threshold_single: float = 1.0
speculative_accept_threshold_acc: float = 1.0
speculative_use_rejection_sampling: bool = False
speculative_token_map: Optional[str] = None
speculative_attention_mode: str = "prefill"
speculative_draft_attention_backend: Optional[str] = None
@@ -6127,6 +6128,12 @@ class ServerArgs:
help="The accept probability of a draft token is raised from its target probability p to min(1, p / threshold_acc).",
default=ServerArgs.speculative_accept_threshold_acc,
)
parser.add_argument(
"--speculative-use-rejection-sampling",
action="store_true",
help="Use rejection sampling for speculative decoding (requires topk=1).",
default=ServerArgs.speculative_use_rejection_sampling,
)
parser.add_argument(
"--speculative-token-map",
type=str,
@@ -32,6 +32,7 @@ from sglang.srt.model_executor.runner_backend.utils import resolve_decode_backen
from sglang.srt.model_executor.runner_backend_utils import (
CUDA_GRAPH_CAPTURE_FAILED_MSG,
)
from sglang.srt.sampling.sampling_batch_info import SamplingBatchInfo
from sglang.srt.speculative.eagle_info import EagleDraftInput
from sglang.srt.utils import (
require_attn_tp_gather,
@@ -59,6 +60,7 @@ class EagleDraftInputBuffers(ForwardInputBuffers):
extend_seq_lens: torch.Tensor
topk_p: torch.Tensor
topk_index: torch.Tensor
draft_probs: Optional[torch.Tensor]
hidden_states: Optional[torch.Tensor]
global_num_tokens_gpu: Optional[torch.Tensor]
global_num_tokens_for_logprob_gpu: Optional[torch.Tensor]
@@ -177,6 +179,14 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
extend_seq_lens = torch.ones((self.max_bs,), dtype=torch.int32)
topk_p = torch.zeros((self.max_bs, self.topk), dtype=torch.float32)
topk_index = torch.zeros((self.max_bs, self.topk), dtype=torch.int64)
draft_probs = (
torch.zeros(
(self.max_bs, self.model_runner.model_config.vocab_size),
dtype=torch.float32,
)
if self.model_runner.server_args.speculative_use_rejection_sampling
else None
)
_hidden_size = EagleDraftInput.hidden_size_for(self.eagle_worker)
hidden_states = (
torch.zeros(
@@ -187,6 +197,8 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
else None
)
self.temperatures = torch.ones((self.max_bs, 1), dtype=torch.float)
if self.require_gathered_buffer:
if self.require_mlp_tp_gather:
global_num_tokens_gpu = torch.zeros(
@@ -222,6 +234,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
extend_seq_lens=extend_seq_lens,
topk_p=topk_p,
topk_index=topk_index,
draft_probs=draft_probs,
hidden_states=hidden_states,
global_num_tokens_gpu=global_num_tokens_gpu,
global_num_tokens_for_logprob_gpu=global_num_tokens_for_logprob_gpu,
@@ -313,6 +326,9 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
)
topk_p = buffers.topk_p[:num_seqs]
topk_index = buffers.topk_index[:num_seqs]
draft_probs = (
buffers.draft_probs[:num_seqs] if buffers.draft_probs is not None else None
)
if self.require_mlp_tp_gather:
global_num_tokens_cpu = [num_tokens] * self.dp_size
@@ -345,10 +361,23 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
spec_info = EagleDraftInput(
topk_p=topk_p,
topk_index=topk_index,
draft_probs=draft_probs,
hidden_states=hidden_states,
capture_hidden_mode=capture_mode,
)
sampling_info = SamplingBatchInfo(
temperatures=self.temperatures[:num_seqs],
top_ps=torch.ones((num_seqs,), dtype=torch.float),
top_ks=torch.full((num_seqs,), -1, dtype=torch.int32),
min_ps=torch.zeros((num_seqs,), dtype=torch.float),
is_all_greedy=False,
need_top_p_sampling=False,
need_top_k_sampling=False,
need_min_p_sampling=False,
vocab_size=self.model_runner.model_config.vocab_size,
)
forward_batch = ForwardBatch(
forward_mode=ForwardMode.DECODE,
batch_size=num_seqs,
@@ -369,6 +398,7 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
global_dp_buffer_len=global_dp_buffer_len,
spec_algorithm=self.model_runner.spec_algorithm,
spec_info=spec_info,
sampling_info=sampling_info,
rids_int=rids_int,
bootstrap_room_ids_int=bootstrap_room_ids_int,
capture_hidden_mode=(
@@ -417,8 +447,10 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
)
def _postprocess_output_to_raw_bs(self, out, raw_bs):
parent_list, top_scores_index, draft_tokens = (t[:raw_bs] for t in out)
return parent_list, top_scores_index, draft_tokens
parent_list, top_scores_index, draft_tokens, draft_probs = (
t[:raw_bs] if t is not None else None for t in out
)
return parent_list, top_scores_index, draft_tokens, draft_probs
# -----------------------------------------------------------------
# Replay
@@ -454,6 +486,8 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
buffers.bootstrap_room_ids_int.fill_(-1)
buffers.topk_p.zero_()
buffers.topk_index.zero_()
if buffers.draft_probs is not None:
buffers.draft_probs.zero_()
if buffers.hidden_states is not None:
buffers.hidden_states.zero_()
buffers.req_pool_indices.zero_()
@@ -505,11 +539,25 @@ class EAGLEDraftCudaGraphRunner(DecodeCudaGraphRunner):
# hidden_states is large + contiguous: copy_() uses the cudaMemcpyAsync
# DMA engine; foreach would force the ~3x slower compute-kernel copy.
if (
buffers.draft_probs is not None
and forward_batch.spec_info.draft_probs is not None
):
buffers.draft_probs[:raw_bs].copy_(forward_batch.spec_info.draft_probs)
if (
buffers.hidden_states is not None
and forward_batch.spec_info.hidden_states is not None
):
buffers.hidden_states[:raw_bs].copy_(forward_batch.spec_info.hidden_states)
# Only rejection sampling reads temperatures (renorm_draft_probs); skip
# the copy otherwise to keep the non-RS path free of extra work.
if (
self.model_runner.server_args.speculative_use_rejection_sampling
and forward_batch.sampling_info is not None
):
self.temperatures[:raw_bs].copy_(
forward_batch.sampling_info.temperatures[:raw_bs]
)
# TODO(ch-wan): support num_token_non_padded
if self.require_gathered_buffer:
@@ -8,6 +8,7 @@ from sglang.srt.constrained.base_grammar_backend import BaseGrammarObject
from sglang.srt.environ import envs
from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
from sglang.srt.server_args import get_global_server_args
from sglang.srt.speculative.eagle_info_v2 import EagleDraftInputV2Mixin
from sglang.srt.speculative.spec_info import SpecInput, SpecInputType
@@ -42,6 +43,9 @@ class EagleVerifyInput(SpecInput):
seq_lens_sum: int
seq_lens_cpu: torch.Tensor
grammar: BaseGrammarObject = None
# Stacked per-step draft proposal distribution q, shape (bs, num_steps,
# vocab); only set under rejection sampling. Consumed by the verify kernel.
draft_probs: torch.Tensor = None
# Shape info for padding
num_tokens_per_req: int = -1 # -1 auto-fills from draft_token_num.
@@ -159,6 +163,9 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
# shape: (b, topk)
topk_p: torch.Tensor = None
topk_index: torch.Tensor = None
# shape: (b, vocab) - single-step draft proposal q from draft-extend;
# only set under rejection sampling.
draft_probs: torch.Tensor = None
# shape: (b, hidden_size) - one hidden per req, consumed by `draft` forward.
# None when the spec algorithm's draft doesn't read hidden_states
# (e.g., STANDALONE — vanilla LLM draft).
@@ -209,6 +216,7 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
dtype: Optional[torch.dtype],
topk: int,
capture_hidden_mode: CaptureHiddenMode,
vocab_size: int = 0,
):
return cls(
bonus_tokens=torch.empty((0,), device=device, dtype=torch.int32),
@@ -219,6 +227,11 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
),
topk_p=torch.empty((0, topk), device=device, dtype=torch.float32),
topk_index=torch.empty((0, topk), device=device, dtype=torch.int64),
draft_probs=(
torch.empty((0, vocab_size), device=device, dtype=torch.float32)
if get_global_server_args().speculative_use_rejection_sampling
else None
),
capture_hidden_mode=capture_hidden_mode,
)
@@ -240,6 +253,8 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
self.topk_p = self.topk_p[: len(new_indices)]
self.topk_index = self.topk_index[: len(new_indices)]
if self.draft_probs is not None:
self.draft_probs = self.draft_probs[: len(new_indices)]
if self.hidden_states is not None:
self.hidden_states = self.hidden_states[: len(new_indices)]
self.bonus_tokens = self.bonus_tokens[: len(new_indices)]
@@ -247,6 +262,8 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
# in some cases(e.g draft_extend), we have not filtered the batch by `unfinished_index`
self.topk_p = self.topk_p[new_indices]
self.topk_index = self.topk_index[new_indices]
if self.draft_probs is not None:
self.draft_probs = self.draft_probs[new_indices]
if self.hidden_states is not None:
self.hidden_states = self.hidden_states[new_indices]
self.bonus_tokens = self.bonus_tokens[new_indices]
@@ -267,6 +284,7 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
self.bonus_tokens = spec_info.bonus_tokens
self.topk_p = spec_info.topk_p
self.topk_index = spec_info.topk_index
self.draft_probs = spec_info.draft_probs
return
if len(spec_info.topk_index) == 0:
return
@@ -279,6 +297,8 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
)
self.topk_p = torch.cat([self.topk_p, spec_info.topk_p])
self.topk_index = torch.cat([self.topk_index, spec_info.topk_index])
if self.draft_probs is not None and spec_info.draft_probs is not None:
self.draft_probs = torch.cat([self.draft_probs, spec_info.draft_probs])
@dataclass
+30 -2
View File
@@ -455,6 +455,14 @@ def eagle_sample(
tree_speculative_sampling_target_only,
)
from sglang.srt.speculative.reject_sampling import (
chain_speculative_sampling_triton,
)
use_rejection_sampling = (
get_global_server_args().speculative_use_rejection_sampling
)
# Apply temperature and get target probs
expanded_temperature = torch.repeat_interleave(
sampling_info.temperatures, verify_input.draft_token_num, dim=0
@@ -479,14 +487,34 @@ def eagle_sample(
)
maybe_detect_nan(target_probs, "v2 verify: target_probs after top_p_renorm")
target_probs = target_probs.reshape(bs, verify_input.draft_token_num, -1)
draft_probs = torch.zeros_like(target_probs)
draft_probs = (
verify_input.draft_probs
if use_rejection_sampling
else torch.zeros_like(target_probs)
)
# Defense-in-depth behind the spec_hook startup allowlist: validate the
# actual kernel inputs (catches draft_probs plumbing regressions or a
# startup guard bypassed by a worker subclass) before the Triton kernel.
if use_rejection_sampling and (
draft_probs is None or draft_probs.shape[-1] != target_probs.shape[-1]
):
raise ValueError(
"Rejection sampling requires a target-vocab draft proposal "
"distribution; the current speculative algorithm/draft worker "
"does not produce one (draft_probs missing or vocab-mismatched)."
)
# coins for rejection sampling
coins = torch.rand_like(candidates, dtype=torch.float32, device=device)
# coins for final sampling
coins_for_final_sampling = torch.rand((bs,), dtype=torch.float32, device=device)
tree_speculative_sampling_target_only(
sampling_fn = (
chain_speculative_sampling_triton
if use_rejection_sampling
else tree_speculative_sampling_target_only
)
sampling_fn(
predicts=predict, # mutable
accept_index=accept_index, # mutable
accept_token_num=num_correct_drafts, # mutable
@@ -74,11 +74,13 @@ from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_utils import (
commit_mamba_states_after_verify,
draft_tp_context,
fast_sample,
generate_token_bitmask,
load_token_map,
move_accept_tokens_to_target_kvcache,
record_stream_each,
record_stream_for_v2_verify,
renorm_draft_probs,
select_top_k_tokens,
spec_stage_span,
)
@@ -146,6 +148,8 @@ class EagleDraftWorker(EagleDraftWorkerBase):
# Args for easy access
self.device = server_args.device
self.topk = server_args.speculative_eagle_topk
if self.server_args.speculative_use_rejection_sampling:
assert self.topk == 1, "Chain speculative sampling supports only topk=1"
self.speculative_num_steps = server_args.speculative_num_steps
self.speculative_num_draft_tokens = server_args.speculative_num_draft_tokens
self.speculative_algorithm = SpeculativeAlgorithm.from_string(
@@ -222,6 +226,22 @@ class EagleDraftWorker(EagleDraftWorkerBase):
self.init_token_map()
self.init_lm_head()
if self.server_args.speculative_use_rejection_sampling:
target_vocab_size = self.target_worker.model_config.vocab_size
draft_vocab_size = (
self.hot_token_id.shape[0]
if self.hot_token_id is not None
else target_vocab_size
)
# FIXME: support reduced (hot) draft vocab by scattering draft probs
# into the target vocab via the d2t map before the sampling kernel.
if draft_vocab_size != target_vocab_size:
raise ValueError(
"--speculative-use-rejection-sampling requires the draft and "
f"target to share one vocab, but the draft vocab "
f"({draft_vocab_size}) != target vocab ({target_vocab_size})."
)
def init_backends(self):
with self.draft_tp_context(
self.draft_runner.tp_group
@@ -443,7 +463,7 @@ class EagleDraftWorker(EagleDraftWorkerBase):
with canary_outside_ctx:
# Run draft
if can_cuda_graph:
parent_list, top_scores_index, draft_tokens = (
parent_list, top_scores_index, draft_tokens, draft_probs = (
self.cuda_graph_runner.execute(forward_batch)
)
else:
@@ -455,8 +475,8 @@ class EagleDraftWorker(EagleDraftWorkerBase):
# `draft_forward` only does sample in this case.
self.draft_attn_backend.init_forward_metadata(forward_batch)
forward_batch.mark_forward_metadata_ready()
parent_list, top_scores_index, draft_tokens = self.draft_forward(
forward_batch
parent_list, top_scores_index, draft_tokens, draft_probs = (
self.draft_forward(forward_batch)
)
if batch.forward_mode.is_idle():
@@ -521,6 +541,7 @@ class EagleDraftWorker(EagleDraftWorkerBase):
capture_hidden_mode=None,
seq_lens_sum=None,
seq_lens_cpu=None,
draft_probs=draft_probs,
)
def draft_forward(self, forward_batch: ForwardBatch):
@@ -549,6 +570,8 @@ class EagleDraftWorker(EagleDraftWorkerBase):
score_list: List[torch.Tensor] = []
token_list: List[torch.Tensor] = []
parents_list: List[torch.Tensor] = []
if self.server_args.speculative_use_rejection_sampling:
draft_probs_list: List[torch.Tensor] = [spec_info.draft_probs]
# Forward multiple steps
scores = None
@@ -598,18 +621,25 @@ class EagleDraftWorker(EagleDraftWorkerBase):
logits_output = self.draft_runner.forward(forward_batch).logits_output
maybe_detect_nan(logits_output.next_token_logits, f"draft_forward step {i}")
maybe_detect_inf(logits_output.next_token_logits, f"draft_forward step {i}")
if self.topk == 1 and not _is_hip:
# topk=1 → degenerate single-path tree; `topk_p` is unused
# downstream, so skip softmax and just argmax over logits.
# Gated to CUDA: on ROCm the argmax tie-break diverges from
# the softmax+max path on FP8 logits and corrupts MTP draft
# selection (DSV3.2 MTP GSM8K, see #26358).
if self.server_args.speculative_use_rejection_sampling:
probs = renorm_draft_probs(
logits_output.next_token_logits,
forward_batch.sampling_info,
self.server_args.speculative_use_rejection_sampling,
)
topk_p, topk_index = fast_sample(probs, num_samples=1)
draft_probs_list.append(probs)
elif self.topk == 1 and not _is_hip:
topk_index = torch.argmax(
logits_output.next_token_logits, dim=-1, keepdim=True
)
topk_p = torch.ones_like(topk_index, dtype=torch.float32)
else:
probs = torch.softmax(logits_output.next_token_logits, dim=-1)
probs = renorm_draft_probs(
logits_output.next_token_logits,
forward_batch.sampling_info,
self.server_args.speculative_use_rejection_sampling,
)
topk_p, topk_index = fast_topk(probs, self.topk, dim=-1)
maybe_detect_oob(
topk_index,
@@ -640,12 +670,24 @@ class EagleDraftWorker(EagleDraftWorkerBase):
draft_tokens = torch.cat(token_list, dim=1)
top_scores_index = self._topk1_score_indices_prealloc[:bs]
parent_list = self._topk1_parents_prealloc[:bs]
return parent_list, top_scores_index, draft_tokens
draft_probs = (
torch.stack(draft_probs_list, dim=1)
if self.server_args.speculative_use_rejection_sampling
else None
)
return parent_list, top_scores_index, draft_tokens, draft_probs
return organize_draft_results(
parent_list, top_scores_index, draft_tokens = organize_draft_results(
score_list, token_list, parents_list, self.speculative_num_draft_tokens
)
draft_probs = (
torch.stack(draft_probs_list, dim=1)
if self.server_args.speculative_use_rejection_sampling
else None
)
return parent_list, top_scores_index, draft_tokens, draft_probs
def draft_extend(self):
pass
@@ -716,11 +758,20 @@ class EagleDraftWorker(EagleDraftWorkerBase):
maybe_detect_inf(logits_output.next_token_logits, "draft_extend_for_prefill")
# Assemble the next-iter draft spec_info from the extend output.
probs = torch.softmax(logits_output.next_token_logits, dim=-1)
topk_p, topk_index = fast_topk(probs, self.topk, dim=-1)
use_rejection_sampling = self.server_args.speculative_use_rejection_sampling
probs = renorm_draft_probs(
logits_output.next_token_logits,
batch.sampling_info,
use_rejection_sampling,
)
if use_rejection_sampling:
topk_p, topk_index = fast_sample(probs, num_samples=1)
else:
topk_p, topk_index = fast_topk(probs, self.topk, dim=-1)
return EagleDraftInput(
topk_p=topk_p,
topk_index=topk_index,
draft_probs=probs if use_rejection_sampling else None,
hidden_states=logits_output.hidden_states,
bonus_tokens=next_token_ids,
num_tokens_per_req=1,
@@ -810,16 +861,30 @@ class EagleDraftWorker(EagleDraftWorkerBase):
]
# The draft-extend graph only anchors full logits; selected-row topk is
# owned by the worker for both graph and eager paths.
if self.topk == 1 and not _is_hip:
if self.server_args.speculative_use_rejection_sampling:
probs = renorm_draft_probs(
draft_logits_output.next_token_logits,
batch.sampling_info,
self.server_args.speculative_use_rejection_sampling,
)
ret_topk_p, ret_topk_index = fast_sample(probs, num_samples=1)
ret_draft_probs = probs
elif self.topk == 1 and not _is_hip:
# Gated to CUDA: see #26358 — ROCm's argmax tie-break corrupts
# MTP draft selection on FP8 logits.
ret_topk_index = torch.argmax(
draft_logits_output.next_token_logits, dim=-1, keepdim=True
)
ret_topk_p = torch.ones_like(ret_topk_index, dtype=torch.float32)
ret_draft_probs = None
else:
probs = torch.softmax(draft_logits_output.next_token_logits, dim=-1)
probs = renorm_draft_probs(
draft_logits_output.next_token_logits,
batch.sampling_info,
self.server_args.speculative_use_rejection_sampling,
)
ret_topk_p, ret_topk_index = fast_topk(probs, self.topk, dim=-1)
ret_draft_probs = None
ret_hidden_states = draft_logits_output.hidden_states
# Construct the return values
@@ -833,6 +898,8 @@ class EagleDraftWorker(EagleDraftWorkerBase):
ret_topk_index,
ret_hidden_states,
)
if self.server_args.speculative_use_rejection_sampling:
next_draft_input.draft_probs = ret_draft_probs
class EAGLEWorkerV2(BaseSpecWorker):
@@ -1010,6 +1077,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
dtype=EagleDraftInput.dtype_for(self.draft_worker),
topk=self.topk,
capture_hidden_mode=capture_mode,
vocab_size=self.target_worker.model_config.vocab_size,
)
if self.speculative_num_steps == 0:
# Drafting disabled (high batch size). _draft_extend below still
@@ -0,0 +1,204 @@
import triton
import triton.language as tl
@triton.jit
def speculative_sampling_classic_kernel(
# Pointers
Predicts,
AcceptIndex,
AcceptTokenNum,
Candidates,
RetriveIndex,
UniformSamples,
UniformSamplesFinal,
TargetProbs,
DraftProbs,
# Strides
stride_cand_b,
stride_cand_s,
stride_idx_b,
stride_idx_s,
stride_uni_b,
stride_uni_s,
stride_tp_b,
stride_tp_s,
stride_tp_v,
stride_dp_b,
stride_dp_s,
stride_dp_v,
# Constants
NUM_SLOTS: tl.constexpr,
VOCAB_SIZE: tl.constexpr,
BLOCK_V: tl.constexpr,
):
pid = tl.program_id(0)
cur_prob_row = 0
cand_ptr_base = Candidates + pid * stride_cand_b
idx_ptr_base = RetriveIndex + pid * stride_idx_b
uni_ptr_base = UniformSamples + pid * stride_uni_b
root_global_idx = tl.load(idx_ptr_base + 0 * stride_idx_s)
tl.store(AcceptIndex + pid * stride_idx_b + 0 * stride_idx_s, root_global_idx)
last_accepted_global_idx = root_global_idx
num_accept = 0
# Verification Loop
step = 1
continue_verifying = 1
while (step < NUM_SLOTS) and (continue_verifying == 1):
draft_token = tl.load(cand_ptr_base + step * stride_cand_s)
offset_prob = (
(pid * stride_tp_b)
+ (cur_prob_row * stride_tp_s)
+ (draft_token * stride_tp_v)
)
offset_draft = (
(pid * stride_dp_b)
+ (cur_prob_row * stride_dp_s)
+ (draft_token * stride_dp_v)
)
p = tl.load(TargetProbs + offset_prob)
q = tl.load(DraftProbs + offset_draft)
coin = tl.load(uni_ptr_base + (step - 1) * stride_uni_s)
if coin * q < p:
num_accept += 1
cur_prob_row = step
tl.store(Predicts + last_accepted_global_idx, draft_token)
curr_global_idx = tl.load(idx_ptr_base + step * stride_idx_s)
tl.store(
AcceptIndex + pid * stride_idx_b + num_accept * stride_idx_s,
curr_global_idx,
)
last_accepted_global_idx = curr_global_idx
step += 1
else:
continue_verifying = 0
tl.store(AcceptTokenNum + pid, num_accept)
# Final Sampling
all_drafts_accepted = continue_verifying
coin_final = tl.load(UniformSamplesFinal + pid)
norm_sum = 0.0
tp_base_ptr = TargetProbs + (pid * stride_tp_b) + (cur_prob_row * stride_tp_s)
# DraftProbs has only num_steps rows (TargetProbs has num_steps + 1). When
# all drafts are accepted cur_prob_row == num_steps is out of bounds for
# DraftProbs, but the all-accepted branch samples pure target p and never
# dereferences this pointer; on rejection cur_prob_row <= num_steps - 1.
dp_base_ptr_safe = DraftProbs + (pid * stride_dp_b) + (cur_prob_row * stride_dp_s)
# Pass 1: Sum
for v_start in range(0, VOCAB_SIZE, BLOCK_V):
v_offsets = v_start + tl.arange(0, BLOCK_V)
mask = v_offsets < VOCAB_SIZE
p_ptr = tp_base_ptr + v_offsets * stride_tp_v
p_val = tl.load(p_ptr, mask=mask, other=0.0)
if all_drafts_accepted:
val = p_val
else:
q_ptr = dp_base_ptr_safe + v_offsets * stride_dp_v
q_val = tl.load(q_ptr, mask=mask, other=0.0)
diff = p_val - q_val
val = tl.where(diff > 0.0, diff, 0.0)
norm_sum += tl.sum(val)
# Pass 2: CDF. Degenerate residual (norm_sum == 0, i.e. p == q everywhere on
# rejection) leaves the cumsum at 0 <= target_u, so final_token falls back to
# VOCAB_SIZE - 1; acceptable since this case is numerically near-impossible.
target_u = coin_final * norm_sum
cum_sum = 0.0
final_token = VOCAB_SIZE - 1
found = 0
for v_start in range(0, VOCAB_SIZE, BLOCK_V):
if found == 0:
v_offsets = v_start + tl.arange(0, BLOCK_V)
mask = v_offsets < VOCAB_SIZE
p_ptr = tp_base_ptr + v_offsets * stride_tp_v
p_val = tl.load(p_ptr, mask=mask, other=0.0)
if all_drafts_accepted:
val = p_val
else:
q_ptr = dp_base_ptr_safe + v_offsets * stride_dp_v
q_val = tl.load(q_ptr, mask=mask, other=0.0)
diff = p_val - q_val
val = tl.where(diff > 0.0, diff, 0.0)
block_cumsum = tl.cumsum(val, axis=0)
total_cumsum = cum_sum + block_cumsum
candidates_mask = total_cumsum > target_u
has_match = tl.max(candidates_mask, axis=0)
if has_match:
match_idx = tl.argmax(candidates_mask.to(tl.int32), axis=0)
final_token = v_start + match_idx
found = 1
cum_sum += tl.sum(val)
tl.store(Predicts + last_accepted_global_idx, final_token)
def chain_speculative_sampling_triton(
predicts,
accept_index,
accept_token_num,
candidates,
retrive_index,
retrive_next_token,
retrive_next_sibling, # not used in chain verification
uniform_samples,
uniform_samples_for_final_sampling,
target_probs,
draft_probs,
threshold_single,
threshold_acc,
deterministic, # not used
):
batch_size, num_slots = candidates.shape
vocab_size = target_probs.shape[-1]
grid = (batch_size,)
speculative_sampling_classic_kernel[grid](
predicts,
accept_index,
accept_token_num,
candidates,
retrive_index,
uniform_samples,
uniform_samples_for_final_sampling,
target_probs,
draft_probs,
candidates.stride(0),
candidates.stride(1),
retrive_index.stride(0),
retrive_index.stride(1),
uniform_samples.stride(0),
uniform_samples.stride(1),
target_probs.stride(0),
target_probs.stride(1),
target_probs.stride(2),
draft_probs.stride(0),
draft_probs.stride(1),
draft_probs.stride(2),
NUM_SLOTS=num_slots,
VOCAB_SIZE=vocab_size,
BLOCK_V=4096,
)
@@ -72,6 +72,28 @@ else:
logger = logging.getLogger(__name__)
def fast_sample(probs: torch.Tensor, num_samples: int = 1):
sample_index = torch.multinomial(probs, num_samples=num_samples)
sample_p = probs.gather(1, sample_index)
return sample_p, sample_index
def renorm_draft_probs(
next_token_logits: torch.Tensor,
sampling_info,
use_rejection_sampling: bool,
) -> torch.Tensor:
"""Draft-side next-token distribution.
Plain softmax, except under rejection sampling where logits are
temperature-scaled so the draft proposal q tracks the target sampling
temperature (higher acceptance; correctness holds for any q).
"""
if not use_rejection_sampling or not next_token_logits.size(0):
return torch.softmax(next_token_logits, dim=-1)
return torch.softmax(next_token_logits / sampling_info.temperatures, dim=-1)
# Simulate acceptance length for benchmarking purposes
SIMULATE_ACC_LEN = envs.SGLANG_SIMULATE_ACC_LEN.get() # turn off if < 0
SIMULATE_ACC_METHOD = envs.SGLANG_SIMULATE_ACC_METHOD.get()