[Spec] Disambiguate verified_id into bonus_token(s) / accept_tokens (#24724)

This commit is contained in:
Liangsheng Yin
2026-05-08 18:24:33 -07:00
committed by GitHub
parent a61a14f416
commit 1613bae412
18 changed files with 128 additions and 108 deletions
@@ -170,7 +170,7 @@ class ScheduleBatchDisaggregationDecodeMixin:
topk_p=topk_p,
topk_index=topk_index,
hidden_states=hidden_states,
verified_id=self.output_ids,
bonus_tokens=self.output_ids,
new_seq_lens=self.seq_lens,
)
spec_info.prepare_for_extend(self)
+3 -3
View File
@@ -362,7 +362,7 @@ def add_output_logprobs_for_spec_v1(
logprobs = torch.nn.functional.log_softmax(
logits_output.next_token_logits / temperatures, dim=-1
)
batch_next_token_ids = res.verified_id
batch_next_token_ids = res.accept_tokens
num_tokens_per_req = [accept + 1 for accept in num_accepted_drafts_per_req_cpu]
# We should repeat top_logprobs_nums to match num_tokens_per_req.
@@ -407,7 +407,7 @@ def add_output_logprobs_for_spec_v1(
# Add output logprobs to the request
pt = 0
next_token_logprobs = logits_output.next_token_logprobs.tolist()
verified_ids = batch_next_token_ids.tolist()
accept_tokens_list = batch_next_token_ids.tolist()
token_top_logprobs_val = logits_output.next_token_top_logprobs_val
token_top_logprobs_idx = logits_output.next_token_top_logprobs_idx
token_ids_logprobs_val = logits_output.next_token_token_ids_logprobs_val
@@ -416,7 +416,7 @@ def add_output_logprobs_for_spec_v1(
for _ in range(num_tokens):
if req.return_logprob:
req.output_token_logprobs_val.append(next_token_logprobs[pt])
req.output_token_logprobs_idx.append(verified_ids[pt])
req.output_token_logprobs_idx.append(accept_tokens_list[pt])
if req.top_logprobs_num > 0:
assert (
should_top_logprobs
+6 -6
View File
@@ -86,7 +86,7 @@ class FutureMap:
# Get a reference for each tensor
topk_p0 = draft_input.topk_p[0]
topk_index0 = draft_input.topk_index[0]
verified_id0 = draft_input.verified_id[0]
bonus_token0 = draft_input.bonus_tokens[0]
new_seq_lens0 = draft_input.new_seq_lens[0]
self.topk_p_buf = torch.empty(
@@ -99,9 +99,9 @@ class FutureMap:
dtype=topk_index0.dtype,
device=self.device,
)
self.verified_id_buf = torch.empty(
(self.future_buffer_len, *verified_id0.shape),
dtype=verified_id0.dtype,
self.bonus_tokens_buf = torch.empty(
(self.future_buffer_len, *bonus_token0.shape),
dtype=bonus_token0.dtype,
device=self.device,
)
self.new_seq_lens_buf = torch.empty(
@@ -146,7 +146,7 @@ class FutureMap:
indices.record_stream(torch.get_device_module(self.device).current_stream())
draft_input.topk_p = self.topk_p_buf[indices]
draft_input.topk_index = self.topk_index_buf[indices]
draft_input.verified_id = self.verified_id_buf[indices]
draft_input.bonus_tokens = self.bonus_tokens_buf[indices]
draft_input.new_seq_lens = self.new_seq_lens_buf[indices]
if spec_need_hidden_states():
draft_input.hidden_states = self.hidden_states_buf[indices]
@@ -181,7 +181,7 @@ class FutureMap:
self.topk_p_buf[intv] = draft_input.topk_p
self.topk_index_buf[intv] = draft_input.topk_index
self.verified_id_buf[intv] = draft_input.verified_id
self.bonus_tokens_buf[intv] = draft_input.bonus_tokens
self.new_seq_lens_buf[intv] = draft_input.new_seq_lens
if spec_need_hidden_states():
self.hidden_states_buf[intv] = draft_input.hidden_states
+13 -11
View File
@@ -67,7 +67,7 @@ class DFlashDraftInput(SpecInput):
"""
# Current token to start the next DFlash block (one per request).
verified_id: torch.Tensor
bonus_tokens: torch.Tensor
# Flattened context features for tokens that need to be appended into the draft cache.
# Shape: [sum(ctx_lens), K * hidden_size], where K is the number of target-layer
@@ -92,7 +92,7 @@ class DFlashDraftInput(SpecInput):
old_ctx_lens = self.ctx_lens
old_target_hidden = self.target_hidden
self.verified_id = self.verified_id[new_indices]
self.bonus_tokens = self.bonus_tokens[new_indices]
self.ctx_lens = old_ctx_lens[new_indices]
self.draft_seq_lens = self.draft_seq_lens[new_indices]
@@ -129,7 +129,9 @@ class DFlashDraftInput(SpecInput):
)
def merge_batch(self, spec_info: "DFlashDraftInput"):
self.verified_id = torch.cat([self.verified_id, spec_info.verified_id], dim=0)
self.bonus_tokens = torch.cat(
[self.bonus_tokens, spec_info.bonus_tokens], dim=0
)
self.ctx_lens = torch.cat([self.ctx_lens, spec_info.ctx_lens], dim=0)
self.draft_seq_lens = torch.cat(
[self.draft_seq_lens, spec_info.draft_seq_lens], dim=0
@@ -319,7 +321,7 @@ class DFlashVerifyInput(SpecInput):
"""DFlash verification for greedy and non-greedy sampling.
Returns:
new_verified_id: int64 tensor [bs] (the new current token per request)
new_bonus_tokens: int64 tensor [bs] (the new current token per request)
commit_lens: int32 tensor [bs] (how many verify-input tokens are committed)
next_target_hidden: tensor [sum(commit_lens), feature_dim]
num_accepted_drafts_per_req_cpu: list[int] (accepted draft tokens per request)
@@ -389,7 +391,7 @@ class DFlashVerifyInput(SpecInput):
max_acc = self.draft_token_num - 1
num_accepted_drafts_per_req_cpu: List[int] = []
commit_lens_cpu: List[int] = []
new_verified_list: List[int] = []
new_bonus_tokens_list: List[int] = []
for i, req in enumerate(batch.reqs):
acc_len = int(packed[i, max_acc].item())
@@ -409,24 +411,24 @@ class DFlashVerifyInput(SpecInput):
req.grammar.accept_token(token_id)
if req.output_ids:
new_verified_token = int(req.output_ids[-1])
new_bonus_token = int(req.output_ids[-1])
elif req.origin_input_ids:
# If no token was appended in this verify step, keep the current token unchanged.
new_verified_token = int(req.origin_input_ids[-1])
new_bonus_token = int(req.origin_input_ids[-1])
else:
raise RuntimeError(
"DFLASH verify cannot determine current token: both output_ids and origin_input_ids are empty."
)
commit_lens_cpu.append(appended)
new_verified_list.append(new_verified_token)
new_bonus_tokens_list.append(new_bonus_token)
num_accepted_drafts_per_req_cpu.append(max(0, appended - 1))
req.spec_verify_ct += 1
req.spec_accepted_drafts += num_accepted_drafts_per_req_cpu[-1]
commit_lens = torch.tensor(commit_lens_cpu, dtype=torch.int32, device=device)
new_verified_id = torch.tensor(
new_verified_list, dtype=torch.int64, device=device
new_bonus_tokens = torch.tensor(
new_bonus_tokens_list, dtype=torch.int64, device=device
)
# Free uncommitted KV cache slots and compact out_cache_loc.
@@ -494,7 +496,7 @@ class DFlashVerifyInput(SpecInput):
logits_output.hidden_states = None
return (
new_verified_id,
new_bonus_tokens,
commit_lens,
next_target_hidden,
num_accepted_drafts_per_req_cpu,
@@ -570,7 +570,7 @@ class DFlashWorker:
block_ids = self._draft_block_ids_buf[:bs]
block_ids.fill_(int(self._mask_token_id))
block_ids[:, 0].copy_(draft_input.verified_id.to(torch.long))
block_ids[:, 0].copy_(draft_input.bonus_tokens.to(torch.long))
noise_embedding = embed_module(block_ids)
input_embeds = noise_embedding.view(-1, noise_embedding.shape[-1])
@@ -1163,7 +1163,7 @@ class DFlashWorker:
model_worker_batch.extend_seq_lens
)
draft_input = DFlashDraftInput(
verified_id=next_token_ids.to(torch.int64),
bonus_tokens=next_token_ids.to(torch.int64),
target_hidden=logits_output.hidden_states,
ctx_lens=extend_seq_lens,
draft_seq_lens=(
@@ -1213,7 +1213,7 @@ class DFlashWorker:
)
(
new_verified_id,
new_bonus_tokens,
commit_lens,
next_target_hidden,
num_accepted_drafts_per_req_cpu,
@@ -1232,7 +1232,7 @@ class DFlashWorker:
# Update draft state for the next iteration. Also materialize the committed verify tokens
# into the draft KV cache immediately so radix cache entries are safe to reuse.
draft_input.verified_id = new_verified_id
draft_input.bonus_tokens = new_bonus_tokens
draft_input.target_hidden = next_target_hidden
draft_input.ctx_lens = commit_lens
self._append_target_hidden_to_draft_kv(batch, draft_input)
@@ -1249,7 +1249,7 @@ class DFlashWorker:
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=new_verified_id,
next_token_ids=new_bonus_tokens,
num_accepted_drafts=num_accepted_drafts,
num_accepted_drafts_per_req_cpu=num_accepted_drafts_per_req_cpu,
can_run_cuda_graph=can_run_cuda_graph,
+35 -19
View File
@@ -250,7 +250,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
capture_hidden_mode=CaptureHiddenMode.LAST,
),
logits_output=logits_output,
verified_id=torch.empty(0, dtype=torch.long, device=batch.device),
accept_tokens=torch.empty(0, dtype=torch.long, device=batch.device),
num_accepted_drafts_per_req_cpu=[],
accepted_indices=torch.full(
(0, self.spec_steps + 1),
@@ -465,7 +465,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
# Free the KV cache for unaccepted tokens
# TODO: fuse them
accept_index = accept_index[accept_index != -1]
verified_id = predict[accept_index]
accept_tokens = predict[accept_index]
evict_mask = torch.full_like(self.draft_token, True, dtype=torch.bool)
evict_mask[accept_index] = False
num_accepted_drafts_cpu = num_accepted_drafts.cpu()
@@ -553,7 +553,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
draft_input = EagleDraftInput(
hidden_states=batch.spec_info.hidden_states[accept_index],
verified_id=verified_id,
accept_tokens=accept_tokens,
num_accepted_drafts=num_accepted_drafts,
num_accepted_tokens=num_accepted_drafts + 1,
num_accepted_drafts_cpu=num_accepted_drafts_list,
@@ -566,7 +566,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
return EagleVerifyOutput(
draft_input=draft_input,
logits_output=logits_output,
verified_id=verified_id,
accept_tokens=accept_tokens,
num_accepted_drafts_per_req_cpu=draft_input.num_accepted_drafts_cpu,
accepted_indices=accept_index,
)
@@ -625,7 +625,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
hidden_states=batch.spec_info.hidden_states[
unfinished_accept_index
],
verified_id=predict[unfinished_accept_index],
accept_tokens=predict[unfinished_accept_index],
num_accepted_drafts_cpu=draft_input_num_accepted_drafts_cpu,
num_accepted_tokens_cpu=draft_input_num_accepted_tokens_cpu,
num_accepted_drafts=unfinished_num_accepted_drafts,
@@ -648,7 +648,7 @@ class EagleVerifyInput(SpecInput, EagleVerifyInputV2Mixin):
return EagleVerifyOutput(
draft_input=draft_input,
logits_output=logits_output,
verified_id=verified_id,
accept_tokens=accept_tokens,
num_accepted_drafts_per_req_cpu=num_accepted_drafts_list,
accepted_indices=accept_index,
)
@@ -669,7 +669,14 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
# `num_accepted_drafts` and `num_accepted_tokens` are kept in sync:
# `num_accepted_tokens = num_accepted_drafts + 1` (per-req, one bonus per req).
# Storing both avoids repeated `+ 1` at every consumer (attn backends, kernels).
verified_id: torch.Tensor = None
bonus_tokens: torch.Tensor = None
# Flat accepted-token tensor for draft-extend, shape `[sum_accepted]`.
# Set right after verify and consumed by `prepare_extend_after_decode` as
# the extend batch's `input_ids`. Dead after that method returns.
# TODO: drop this field and pass `accept_tokens` directly to
# `prepare_extend_after_decode` as a method arg. Its lifetime is bounded
# by verify -> prepare_extend, no need to live on the dataclass.
accept_tokens: torch.Tensor = None
num_accepted_drafts: torch.Tensor = None
num_accepted_tokens: torch.Tensor = None
num_accepted_drafts_cpu: List[int] = None
@@ -707,13 +714,13 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
return
# Prefill only generate 1 token.
assert len(self.verified_id) == len(batch.seq_lens)
assert len(self.bonus_tokens) == len(batch.seq_lens)
pt = 0
for i, extend_len in enumerate(batch.extend_lens):
input_ids = batch.input_ids[pt : pt + extend_len]
batch.input_ids[pt : pt + extend_len] = torch.cat(
(input_ids[1:], self.verified_id[i].reshape(1))
(input_ids[1:], self.bonus_tokens[i].reshape(1))
)
pt += extend_len
@@ -727,7 +734,8 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
capture_hidden_mode: CaptureHiddenMode,
):
return cls(
verified_id=torch.empty((0,), device=device, dtype=torch.int32),
bonus_tokens=torch.empty((0,), device=device, dtype=torch.int32),
accept_tokens=torch.empty((0,), device=device, dtype=torch.int32),
hidden_states=torch.empty((0, hidden_size), device=device, dtype=dtype),
topk_p=torch.empty((0, topk), device=device, dtype=torch.float32),
topk_index=torch.empty((0, topk), device=device, dtype=torch.int64),
@@ -748,7 +756,11 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
if batch.forward_mode.is_idle():
return
batch.input_ids = self.verified_id
# `self.accept_tokens` is the flat accepted-token tensor set by
# `EagleVerifyInput.verify`; use it as the extend batch's `input_ids`.
# The kernel below populates `self.bonus_tokens` ([bs] per-req) for
# the next decode round.
batch.input_ids = self.accept_tokens
batch.extend_lens = batch.spec_info.num_accepted_tokens_cpu
batch.extend_num_tokens = sum(batch.extend_lens)
batch.seq_lens = batch.spec_info.seq_lens_for_draft_extend
@@ -759,14 +771,16 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
self.capture_hidden_mode = CaptureHiddenMode.LAST
self.positions = torch.empty_like(batch.input_ids, dtype=torch.long)
self.verified_id = torch.empty_like(self.num_accepted_tokens, dtype=torch.int32)
self.bonus_tokens = torch.empty_like(
self.num_accepted_tokens, dtype=torch.int32
)
create_extend_after_decode_spec_info[(len(batch.seq_lens),)](
batch.input_ids,
batch.seq_lens,
self.num_accepted_tokens,
self.positions,
self.verified_id,
self.bonus_tokens,
next_power_of_2(max(speculative_num_steps + 1, len(batch.seq_lens))),
)
@@ -821,13 +835,13 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
self.topk_p = self.topk_p[: len(new_indices)]
self.topk_index = self.topk_index[: len(new_indices)]
self.hidden_states = self.hidden_states[: len(new_indices)]
self.verified_id = self.verified_id[: len(new_indices)]
self.bonus_tokens = self.bonus_tokens[: len(new_indices)]
else:
# 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]
self.hidden_states = self.hidden_states[new_indices]
self.verified_id = self.verified_id[new_indices]
self.bonus_tokens = self.bonus_tokens[new_indices]
def merge_batch(self, spec_info: "EagleDraftInput"):
if self.future_indices is not None:
@@ -841,7 +855,7 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
if self.hidden_states is None:
self.hidden_states = spec_info.hidden_states
self.verified_id = spec_info.verified_id
self.bonus_tokens = spec_info.bonus_tokens
self.topk_p = spec_info.topk_p
self.topk_index = spec_info.topk_index
return
@@ -850,7 +864,9 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
self.hidden_states = torch.cat(
[self.hidden_states, spec_info.hidden_states], axis=0
)
self.verified_id = torch.cat([self.verified_id, spec_info.verified_id], axis=0)
self.bonus_tokens = torch.cat(
[self.bonus_tokens, spec_info.bonus_tokens], axis=0
)
self.topk_p = torch.cat([self.topk_p, spec_info.topk_p])
self.topk_index = torch.cat([self.topk_index, spec_info.topk_index])
@@ -861,8 +877,8 @@ class EagleVerifyOutput:
draft_input: EagleDraftInput
# Logit outputs from target worker
logits_output: LogitsProcessorOutput
# Accepted token ids including the bonus token
verified_id: torch.Tensor
# Accepted token ids including the bonus token (flat, [sum_accepted])
accept_tokens: torch.Tensor
# Accepted token length per sequence in a batch in CPU.
num_accepted_drafts_per_req_cpu: List[int]
# Accepted indices from logits_output.next_token_logits
@@ -473,10 +473,10 @@ class EagleVerifyInputV2Mixin:
@triton.jit
def fill_new_verified_id(
verified_id,
def fill_bonus_tokens(
accept_tokens,
accept_lens,
new_verified_id,
bonus_tokens_ptr,
num_draft_tokens: tl.constexpr,
):
# NOTE: we cannot fuse any in-place operations of `accept_lens` inside this kernel
@@ -485,9 +485,9 @@ def fill_new_verified_id(
# `accept_lens` includes the bonus token; the last accepted slot is at -1.
accept_len = tl.load(accept_lens + pid)
verified_id_idx = num_draft_tokens * pid + accept_len - 1
verified_id_data = tl.load(verified_id + verified_id_idx)
tl.store(new_verified_id + pid, verified_id_data)
bonus_token_idx = num_draft_tokens * pid + accept_len - 1
bonus_token = tl.load(accept_tokens + bonus_token_idx)
tl.store(bonus_tokens_ptr + pid, bonus_token)
@triton.jit
+2 -2
View File
@@ -46,7 +46,7 @@ class TreeMaskMode(IntEnum):
def build_tree_kernel_efficient(
verified_id: torch.Tensor,
bonus_tokens: torch.Tensor,
parent_list: List[torch.Tensor],
top_scores_index: torch.Tensor,
draft_tokens: torch.Tensor,
@@ -59,7 +59,7 @@ def build_tree_kernel_efficient(
tree_mask_buf: Optional[torch.Tensor] = None,
position_buf: Optional[torch.Tensor] = None,
):
draft_tokens = torch.cat((verified_id.unsqueeze(1), draft_tokens), dim=1).flatten()
draft_tokens = torch.cat((bonus_tokens.unsqueeze(1), draft_tokens), dim=1).flatten()
# seq_lens_sum == sum(seq_lens); seq_lens: sequence length without draft tokens
bs = seq_lens.numel()
@@ -505,7 +505,7 @@ class EAGLEWorker(TpModelWorker):
# when DP attention is enabled, but it is slow. Skip it for now.
if (
self.server_args.enable_dp_attention
or batch.spec_info.verified_id.shape[0] > 0
or batch.spec_info.accept_tokens.shape[0] > 0
):
# decode is not finished
self.forward_draft_extend_after_decode(batch)
@@ -521,14 +521,14 @@ class EAGLEWorker(TpModelWorker):
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=verify_output.verified_id,
next_token_ids=verify_output.accept_tokens,
num_accepted_drafts=sum(verify_output.num_accepted_drafts_per_req_cpu),
num_accepted_drafts_per_req_cpu=verify_output.num_accepted_drafts_per_req_cpu,
can_run_cuda_graph=can_run_cuda_graph,
)
def check_forward_draft_extend_after_decode(self, batch: ScheduleBatch):
local_need_forward = batch.spec_info.verified_id.shape[0] > 0
local_need_forward = batch.spec_info.accept_tokens.shape[0] > 0
if not self.server_args.enable_dp_attention:
return local_need_forward
@@ -588,7 +588,7 @@ class EAGLEWorker(TpModelWorker):
if batch.sampling_info.penalizer_orchestrator.is_required:
# This is a relaxed version of penalties for speculative decoding.
batch.sampling_info.penalizer_orchestrator.cumulate_output_tokens(
spec_info.verified_id.to(torch.int64)
spec_info.bonus_tokens.to(torch.int64)
)
# Allocate cache locations
@@ -779,7 +779,7 @@ class EAGLEWorker(TpModelWorker):
retrieve_next_sibling,
draft_tokens,
) = build_tree_kernel_efficient(
spec_info.verified_id,
spec_info.bonus_tokens,
parent_list,
top_scores_index,
draft_tokens,
@@ -1082,7 +1082,7 @@ class EAGLEWorker(TpModelWorker):
"""
batch.spec_info = EagleDraftInput(
hidden_states=hidden_states,
verified_id=next_token_ids,
bonus_tokens=next_token_ids,
num_tokens_per_req=1,
num_tokens_for_logprob_per_req=1,
)
@@ -1116,7 +1116,7 @@ class EAGLEWorker(TpModelWorker):
input_is_idle = batch.forward_mode.is_idle()
if not input_is_idle and batch.spec_info.verified_id.numel() == 0:
if not input_is_idle and batch.spec_info.accept_tokens.numel() == 0:
batch = batch.copy()
batch.prepare_for_idle()
hidden_size = (
@@ -49,7 +49,7 @@ from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput
from sglang.srt.speculative.eagle_info_v2 import (
assign_extend_cache_locs,
fill_accepted_out_cache_loc,
fill_new_verified_id,
fill_bonus_tokens,
)
from sglang.srt.speculative.eagle_utils import TreeMaskMode, build_tree_kernel_efficient
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
@@ -383,7 +383,7 @@ class EagleDraftWorker(BaseDraftWorker):
retrieve_next_sibling,
draft_tokens,
) = build_tree_kernel_efficient(
draft_input.verified_id,
draft_input.bonus_tokens,
parent_list,
top_scores_index,
draft_tokens,
@@ -537,7 +537,7 @@ class EagleDraftWorker(BaseDraftWorker):
# Construct spec_info
next_draft_input = EagleDraftInput(
hidden_states=target_hidden_states,
verified_id=next_token_ids,
bonus_tokens=next_token_ids,
new_seq_lens=batch.seq_lens,
# draft mode is same with decode mode, only 1 token per req
num_tokens_per_req=1,
@@ -1041,16 +1041,16 @@ class EAGLEWorkerV2(BaseSpecWorker):
verify_done.record()
if not batch.forward_mode.is_idle():
all_verified_id = predict[accept_index]
verified_id = torch.empty_like(accept_lens, dtype=torch.int32)
fill_new_verified_id[(bs,)](
all_verified_id,
accept_tokens = predict[accept_index]
bonus_tokens = torch.empty_like(accept_lens, dtype=torch.int32)
fill_bonus_tokens[(bs,)](
accept_tokens,
accept_lens,
verified_id,
bonus_tokens,
self.speculative_num_draft_tokens,
)
else:
verified_id = torch.empty((0,), device=self.device, dtype=torch.int32)
bonus_tokens = torch.empty((0,), device=self.device, dtype=torch.int32)
if batch.return_logprob and not batch.forward_mode.is_idle():
compute_spec_v2_logprobs(
@@ -1059,7 +1059,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
# Construct the next draft input
next_draft_input = EagleDraftInput(
verified_id=verified_id,
bonus_tokens=bonus_tokens,
new_seq_lens=new_seq_lens,
verify_done=verify_done,
)
@@ -137,12 +137,12 @@ def select_last_verified_seed(
draft_input: FrozenKVMTPDraftInput,
) -> Tuple[torch.Tensor, torch.Tensor]:
if draft_input.num_accepted_tokens is None:
return draft_input.verified_id, draft_input.hidden_states
return draft_input.bonus_tokens, draft_input.hidden_states
counts = draft_input.num_accepted_tokens.to(torch.long)
last_indices = torch.cumsum(counts, dim=0) - 1
return (
draft_input.verified_id[last_indices],
draft_input.bonus_tokens[last_indices],
draft_input.hidden_states[last_indices],
)
@@ -366,7 +366,7 @@ class FrozenKVMTPWorker(TpModelWorker):
if draft_input is None:
draft_input = FrozenKVMTPDraftInput()
draft_input.verified_id = last_token_ids.to(torch.int64)
draft_input.bonus_tokens = last_token_ids.to(torch.int64)
draft_input.hidden_states = last_hidden_states
draft_input.capture_hidden_mode = CaptureHiddenMode.LAST
draft_input.num_tokens_per_req = 1
@@ -380,7 +380,7 @@ class FrozenKVMTPWorker(TpModelWorker):
spec_info_backup = batch.spec_info
batch.forward_mode = ForwardMode.DECODE
batch.input_ids = draft_input.verified_id
batch.input_ids = draft_input.bonus_tokens
batch.return_hidden_states = False
batch.return_logprob = False
batch.spec_info = draft_input
@@ -461,14 +461,14 @@ class FrozenKVMTPWorker(TpModelWorker):
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
if (
self.server_args.enable_dp_attention
or batch.spec_info.verified_id.numel()
or batch.spec_info.bonus_tokens.numel()
):
self.forward_draft_extend_after_decode(batch)
set_time_batch(batch.reqs, "set_spec_draft_extend_end_time", trace_only=True)
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=verify_output.verified_id,
next_token_ids=verify_output.accept_tokens,
num_accepted_drafts=sum(verify_output.num_accepted_drafts_per_req_cpu),
num_accepted_drafts_per_req_cpu=verify_output.num_accepted_drafts_per_req_cpu,
can_run_cuda_graph=can_run_cuda_graph,
@@ -507,7 +507,7 @@ class FrozenKVMTPWorker(TpModelWorker):
def forward_draft_extend_after_decode(self, batch: ScheduleBatch) -> None:
assert isinstance(batch.spec_info, FrozenKVMTPDraftInput)
input_is_idle = batch.forward_mode.is_idle()
if not input_is_idle and batch.spec_info.verified_id.numel() == 0:
if not input_is_idle and batch.spec_info.bonus_tokens.numel() == 0:
batch = batch.copy()
batch.prepare_for_idle()
batch.spec_info = FrozenKVMTPDraftInput.create_idle_input(
@@ -564,7 +564,7 @@ class FrozenKVMTPWorker(TpModelWorker):
if batch.sampling_info.penalizer_orchestrator.is_required:
batch.sampling_info.penalizer_orchestrator.cumulate_output_tokens(
spec_info.verified_id.to(torch.int64)
spec_info.bonus_tokens.to(torch.int64)
)
spec_info.capture_hidden_mode = CaptureHiddenMode.LAST
@@ -603,7 +603,7 @@ class FrozenKVMTPWorker(TpModelWorker):
retrieve_next_sibling,
draft_tokens,
) = build_tree_kernel_efficient(
spec_info.verified_id,
spec_info.bonus_tokens,
parent_list,
top_scores_index,
draft_tokens,
@@ -283,20 +283,20 @@ class MultiLayerEagleWorker(TpModelWorker):
# when DP attention is enabled, but it is slow. Skip it for now.
if (
self.server_args.enable_dp_attention
or batch.spec_info.verified_id.shape[0] > 0
or batch.spec_info.accept_tokens.shape[0] > 0
):
# decode is not finished
self.forward_draft_extend_after_decode(batch)
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=verify_output.verified_id,
next_token_ids=verify_output.accept_tokens,
num_accepted_drafts=sum(verify_output.num_accepted_drafts_per_req_cpu),
can_run_cuda_graph=can_run_cuda_graph,
)
def check_forward_draft_extend_after_decode(self, batch: ScheduleBatch):
local_need_forward = batch.spec_info.verified_id.shape[0] > 0
local_need_forward = batch.spec_info.accept_tokens.shape[0] > 0
if not self.server_args.enable_dp_attention:
return local_need_forward
@@ -440,7 +440,7 @@ class MultiLayerEagleWorker(TpModelWorker):
retrieve_next_sibling,
draft_tokens,
) = build_tree_kernel_efficient(
spec_info.verified_id,
spec_info.bonus_tokens,
parent_list,
top_scores_index,
draft_tokens,
@@ -609,7 +609,7 @@ class MultiLayerEagleWorker(TpModelWorker):
"""
batch.spec_info = EagleDraftInput(
hidden_states=hidden_states,
verified_id=next_token_ids,
bonus_tokens=next_token_ids,
num_tokens_per_req=1,
num_tokens_for_logprob_per_req=1,
)
@@ -664,7 +664,7 @@ class MultiLayerEagleWorker(TpModelWorker):
input_is_idle = batch.forward_mode.is_idle()
if not input_is_idle and batch.spec_info.verified_id.numel() == 0:
if not input_is_idle and batch.spec_info.accept_tokens.numel() == 0:
batch = batch.copy()
batch.prepare_for_idle()
hidden_size = (
@@ -33,7 +33,7 @@ from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.base_spec_worker import BaseDraftWorker, BaseSpecWorker
from sglang.srt.speculative.draft_utils import DraftBackendFactory
from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput
from sglang.srt.speculative.eagle_info_v2 import fill_new_verified_id
from sglang.srt.speculative.eagle_info_v2 import fill_bonus_tokens
from sglang.srt.speculative.eagle_utils import TreeMaskMode, build_tree_kernel_efficient
from sglang.srt.speculative.multi_layer_eagle_draft_extend_cuda_graph_runner import (
MultiLayerEagleMultiStepDraftExtendCudaGraphRunner,
@@ -254,7 +254,7 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
retrieve_next_sibling,
draft_tokens,
) = build_tree_kernel_efficient(
draft_input.verified_id,
draft_input.bonus_tokens,
parent_list,
top_scores_index,
draft_tokens,
@@ -374,7 +374,7 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
# Construct spec_info
next_draft_input = EagleDraftInput(
hidden_states=target_hidden_states,
verified_id=next_token_ids,
bonus_tokens=next_token_ids,
new_seq_lens=batch.seq_lens,
# draft mode is same with decode mode, only 1 token per req
num_tokens_per_req=1,
@@ -762,16 +762,16 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
verify_done.record()
if not batch.forward_mode.is_idle():
all_verified_id = predict[accept_index]
verified_id = torch.empty_like(accept_lens, dtype=torch.int32)
fill_new_verified_id[(bs,)](
all_verified_id,
accept_tokens = predict[accept_index]
bonus_tokens = torch.empty_like(accept_lens, dtype=torch.int32)
fill_bonus_tokens[(bs,)](
accept_tokens,
accept_lens,
verified_id,
bonus_tokens,
self.speculative_num_draft_tokens,
)
else:
verified_id = torch.empty((0,), device=self.device, dtype=torch.int32)
bonus_tokens = torch.empty((0,), device=self.device, dtype=torch.int32)
if batch.return_logprob and not batch.forward_mode.is_idle():
compute_spec_v2_logprobs(
@@ -780,7 +780,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
# Construct the next draft input
next_draft_input = EagleDraftInput(
verified_id=verified_id,
bonus_tokens=bonus_tokens,
new_seq_lens=new_seq_lens,
verify_done=verify_done,
)
+2 -2
View File
@@ -206,7 +206,7 @@ class NgramVerifyInput(SpecInput):
logits_output.hidden_states = logits_output.hidden_states[
self.accepted_indices
]
self.verified_id = self.predict[self.accepted_indices]
self.accept_tokens = self.predict[self.accepted_indices]
def _free_cache(
self,
@@ -465,7 +465,7 @@ class NgramVerifyInput(SpecInput):
batch.seq_lens.add_(self.num_accepted_tokens)
batch.seq_lens_cpu.add_(num_accepted_tokens_cpu)
return logits_output, self.verified_id, num_accepted_drafts
return logits_output, self.accept_tokens, num_accepted_drafts
def filter_batch(self, new_indices: torch.Tensor, has_been_filtered: bool = True):
pass
+4 -4
View File
@@ -62,11 +62,11 @@ def spec_need_hidden_states(server_args: Optional[ServerArgs] = None) -> bool:
@triton.jit
def create_extend_after_decode_spec_info(
verified_id,
accept_tokens,
seq_lens,
accept_lens,
positions,
new_verified_id,
bonus_tokens_ptr,
bs_upper: tl.constexpr,
):
pid = tl.program_id(axis=0)
@@ -83,8 +83,8 @@ def create_extend_after_decode_spec_info(
tl.store(positions_ptr + offsets, seq_length - accept_len + offsets, mask)
accept_len_cumsum += accept_len - 1
verified_id_data = tl.load(verified_id + accept_len_cumsum)
tl.store(new_verified_id + pid, verified_id_data)
bonus_token = tl.load(accept_tokens + accept_len_cumsum)
tl.store(bonus_tokens_ptr + pid, bonus_token)
@triton.jit