[Spec] Consolidate the verify step into eagle_worker_common.run_eagle_verify (#31380)

This commit is contained in:
Liangsheng Yin
2026-07-16 00:56:54 -07:00
committed by GitHub
parent 5cbea10e2f
commit b55228cfdb
3 changed files with 311 additions and 326 deletions
@@ -7,21 +7,38 @@ import torch
from sglang.kernels.ops.speculative.cache_locs import (
assign_draft_cache_locs_contiguous,
)
from sglang.kernels.ops.speculative.eagle import fill_bonus_tokens_func
from sglang.srt.layers.utils.logprob import compute_spec_v2_logprobs
from sglang.srt.managers.utils import GenerationBatchResult
from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode,
ForwardBatch,
ForwardMode,
)
from sglang.srt.speculative.eagle_info import EagleVerifyInput
from sglang.srt.speculative.eagle_info import EagleDraftInput, EagleVerifyInput
from sglang.srt.speculative.eagle_utils import (
TreeMaskMode,
build_tree_kernel_efficient,
eagle_prepare_for_verify,
eagle_sample,
)
from sglang.srt.speculative.spec_utils import (
commit_mamba_states_after_verify,
generate_token_bitmask,
move_accept_tokens_to_target_kvcache,
record_stream_each,
record_stream_for_v2_verify,
)
from sglang.srt.utils import is_cpu
from sglang.srt.utils.async_probe import maybe_detect_oob
from sglang.srt.utils.async_probe import (
maybe_detect_inf,
maybe_detect_nan,
maybe_detect_oob,
)
from sglang.srt.utils.common import is_npu
_is_cpu = is_cpu()
_is_npu = is_npu()
if _is_cpu:
from sgl_kernel import assign_draft_cache_locs_contiguous_cpu
@@ -34,10 +51,7 @@ if TYPE_CHECKING:
from sglang.srt.speculative.eagle_draft_cuda_graph_runner import (
EAGLEDraftCudaGraphRunner,
)
from sglang.srt.speculative.eagle_info import (
EagleDraftExtendInput,
EagleDraftInput,
)
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
def duplicate_prefix_tail_to_draft_branches(
@@ -357,3 +371,266 @@ def build_eagle_verify_input(
seq_lens_cpu=None,
draft_probs=draft_probs,
)
def _finalize_accept_tree_path(
batch: ScheduleBatch,
accept_index: torch.Tensor,
accept_lens: torch.Tensor,
predict: torch.Tensor,
logits_output: Any,
bs: int,
*,
token_to_kv_pool_allocator: Any,
num_draft_tokens: int,
) -> torch.Tensor:
"""Tree drafting (topk > 1): move the accepted path -- KV slots, predict,
hidden_states -- to the contiguous front of each per-req block, which the
downstream chain-layout code (draft-extend select_index, committed-KV reads)
assumes. Returns compacted predict; mutates logits_output.hidden_states
(moved only when present)."""
move_accept_tokens_to_target_kvcache(
batch, accept_index, accept_lens - 1, token_to_kv_pool_allocator
)
predict = _compact_accept_to_front(
predict, accept_index, bs, num_draft_tokens=num_draft_tokens
)
if logits_output.hidden_states is not None:
logits_output.hidden_states = _compact_accept_to_front(
logits_output.hidden_states,
accept_index,
bs,
num_draft_tokens=num_draft_tokens,
)
return predict
def _compact_accept_to_front(
x: torch.Tensor,
accept_index: torch.Tensor,
bs: int,
*,
num_draft_tokens: int,
) -> torch.Tensor:
"""Gather the accepted tree path to the front of each per-req block.
``x`` is node-indexed over the whole tree (``[bs * num_draft_tokens, ...]``),
``accept_index`` is ``[bs, spec_steps + 1]`` global node indices (-1 padded).
Padded entries clamp to node 0 but land past accept_lens (never read);
trailing unaccepted slots stay and are freed as overshoot.
"""
nd = num_draft_tokens
s1 = accept_index.shape[1] # spec_steps + 1
safe = accept_index.to(torch.int64).clamp(min=0).reshape(-1)
gathered = x[safe]
out = x.clone()
out.view(bs, nd, *x.shape[1:])[:, :s1] = gathered.view(bs, s1, *x.shape[1:])
return out
def run_eagle_verify(
batch: ScheduleBatch,
*,
target_worker: TpModelWorker,
req_to_token_pool: ReqToTokenPool,
token_to_kv_pool_allocator: Any,
plan_stream: Any,
plan_stream_ctx: Any,
topk: int,
num_steps: int,
num_draft_tokens: int,
device: str,
metadata_ready_pre_pad: bool,
finalize_tree_path: bool,
) -> GenerationBatchResult:
"""Shared verify step: target-verify forward, sampling, acceptance bookkeeping.
The single-layer eagle verify body is the source of truth (superset). Two
switches encode the multi-layer worker's preserved-verbatim differences:
- ``metadata_ready_pre_pad``: multi-layer marks forward metadata ready
pre-pad unconditionally; single-layer relies on eagle_prepare_for_verify
marking it only when the cuda-graph path ran.
- ``finalize_tree_path``: single-layer compacts the accepted tree path to
the front of each per-req block for topk > 1; multi-layer has never run
this compaction.
"""
fwd_stream = torch.get_device_module(device).current_stream()
verify_input: EagleVerifyInput = batch.spec_info
record_stream_for_v2_verify(batch, verify_input, fwd_stream)
bs = len(batch.seq_lens)
# Batch 1: Target verify
# Prepare for target verify in a separate stream
with plan_stream_ctx:
verify_forward_batch, can_run_cuda_graph = eagle_prepare_for_verify(
verify_input,
req_to_token_pool,
batch,
target_worker,
)
# Cover post-prepare rebinds: draft_token, plan_stream-allocated out_cache_loc.
record_stream_each((batch.input_ids, batch.out_cache_loc), fwd_stream)
# Correct some buffers due to the overlap plan
if plan_stream:
torch.get_device_module(device).current_stream().wait_stream(plan_stream)
if (
_is_npu
and target_worker.model_runner.model_config.model_is_mrope
and batch.spec_info is not None
and getattr(batch.spec_info, "positions", None) is not None
and not batch.forward_mode.is_idle()
):
# mrope_position depends on draft output in default stream and is computed in plan stream,
# causing errors. Compute it here for correct values.
verify_forward_batch.compute_spec_mrope_positions(
target_worker.model_runner, batch
)
# Some values such as custom_mask and position depend on the output of draft,
# so the previous plan step used the wrong values. Here, we need to run the related
# computation again to update them to the correct values.
target_worker.model_runner.attn_backend.update_verify_buffers_to_fill_after_draft(
verify_input,
(
target_worker.model_runner.decode_cuda_graph_runner.bs
if can_run_cuda_graph
else None
),
)
# Prepare grammar data on CPU if needed
if batch.has_grammar:
retrieve_next_token_cpu = verify_input.retrieve_next_token.cpu()
retrieve_next_sibling_cpu = verify_input.retrieve_next_sibling.cpu()
draft_tokens_cpu = verify_input.draft_token.view(
verify_input.retrieve_next_token.shape
).cpu()
if metadata_ready_pre_pad:
# Multi-layer eagle preserved-verbatim behavior: metadata init is
# skipped here unconditionally, although eagle_prepare_for_verify
# only plans when cuda-graph load_batch ran. Single-layer eagle
# re-inits the non-graph path instead (post-pad); multi-layer has
# not adopted that fix. On NPU with --disable-cuda-graph, non-graph
# verify needs metadata init in forward_extend (post-pad); only
# mark ready for the cuda-graph path.
if not _is_npu or can_run_cuda_graph:
verify_forward_batch.mark_forward_metadata_ready()
# Run target verify batch in the main compute stream (GPU compute).
# Metadata init is skipped iff cuda-graph already ran load_batch —
# eagle_prepare_for_verify marked the batch in exactly that case; the
# non-cuda-graph path stays unmarked and gets forward_extend's init
# (post-pad).
forward_batch_output = target_worker.forward_batch_generation(
batch=None,
forward_batch=verify_forward_batch,
is_verify=True,
)
logits_output = forward_batch_output.logits_output
# Generate vocab mask for constrained decoding
vocab_mask = None
if batch.has_grammar:
# Generate the logit mask for structured output.
vocab_mask = generate_token_bitmask(
batch.reqs,
verify_input,
retrieve_next_token_cpu,
retrieve_next_sibling_cpu,
draft_tokens_cpu,
batch.sampling_info.vocab_size,
)
if vocab_mask is not None:
assert verify_input.grammar is not None
vocab_mask = vocab_mask.to(verify_input.retrieve_next_token.device)
# NOTE: otherwise, this vocab mask will be the one from the previous extend stage
# and will be applied to produce wrong results
batch.sampling_info.vocab_mask = None
# Sample
maybe_detect_nan(logits_output.next_token_logits, "verify: target model logits")
maybe_detect_inf(logits_output.next_token_logits, "verify: target model logits")
(
predict,
accept_lens,
accept_index,
) = eagle_sample(verify_input, batch, logits_output, vocab_mask)
new_seq_lens = batch.seq_lens + accept_lens
clear_unaccepted_c128 = getattr(
token_to_kv_pool_allocator.get_kvcache(),
"clear_unaccepted_c128_draft_states",
None,
)
if clear_unaccepted_c128 is not None and not batch.forward_mode.is_idle():
clear_unaccepted_c128(
batch.req_pool_indices,
batch.seq_lens,
accept_lens,
num_draft_tokens,
)
# Update mamba state for hybrid GDN models after verification
commit_mamba_states_after_verify(
target_worker,
batch,
accept_lens,
accept_index,
num_draft_tokens,
)
if not batch.forward_mode.is_idle():
accept_tokens = predict[accept_index]
bonus_tokens = torch.empty_like(accept_lens, dtype=torch.int32)
# stride = accept_tokens per-req width = accept_index.shape[1]
# (spec_steps + 1); NOT num_draft_tokens, wrong for topk > 1 trees.
fill_bonus_tokens_func(
accept_tokens,
accept_lens,
bonus_tokens,
accept_index.shape[1],
bs,
)
else:
bonus_tokens = torch.empty((0,), device=device, dtype=torch.int32)
if batch.return_logprob and not batch.forward_mode.is_idle():
compute_spec_v2_logprobs(batch, logits_output, predict, accept_index, num_steps)
if finalize_tree_path and not batch.forward_mode.is_idle() and topk > 1:
# topk == 1 needs nothing here: the accepted path is already the front
# chain, so the whole compaction is an identity transform.
predict = _finalize_accept_tree_path(
batch,
accept_index,
accept_lens,
predict,
logits_output,
bs,
token_to_kv_pool_allocator=token_to_kv_pool_allocator,
num_draft_tokens=num_draft_tokens,
)
next_draft_input = EagleDraftInput(bonus_tokens=bonus_tokens)
# verify_forward_batch transitively holds verify-time GPU tensors
# (draft_token / out_cache_loc / ...) that must outlive the imminent
# batch.input_ids rebind in prepare_for_draft_extend.
# Scheduler pins it in batch_record_buf for the 2-iter window.
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=predict,
can_run_cuda_graph=can_run_cuda_graph,
speculative_num_draft_tokens=num_draft_tokens,
next_draft_input=next_draft_input,
accept_lens=accept_lens,
new_seq_lens=new_seq_lens,
routed_experts_output=forward_batch_output.routed_experts_output,
indexer_topk_output=forward_batch_output.indexer_topk_output,
extra_keep_alive_refs=[verify_forward_batch],
)
+13 -214
View File
@@ -6,7 +6,6 @@ from typing import List, Optional
import torch
from sglang.kernels.ops.speculative.eagle import fill_bonus_tokens_func
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs
from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_extend_npu_graph_runner import (
@@ -28,7 +27,6 @@ from sglang.srt.layers.moe.utils import (
speculative_moe_a2a_backend_context,
speculative_moe_backend_context,
)
from sglang.srt.layers.utils.logprob import compute_spec_v2_logprobs
from sglang.srt.managers.io_struct import UpdateWeightsFromTensorReqInput
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.scheduler import GenerationBatchResult
@@ -66,8 +64,6 @@ from sglang.srt.speculative.eagle_info import (
from sglang.srt.speculative.eagle_utils import (
_eagle_prefill_tail_tokens,
default_tree_mask_mode,
eagle_prepare_for_verify,
eagle_sample,
get_draft_recurrent_hidden_state_spec,
organize_draft_results,
per_step_draft_out_cache_loc,
@@ -76,18 +72,14 @@ from sglang.srt.speculative.eagle_worker_common import (
build_eagle_verify_input,
prepare_for_draft,
prepare_for_draft_extend,
run_eagle_verify,
)
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,
get_plan_stream,
load_token_map,
move_accept_tokens_to_target_kvcache,
record_stream_each,
record_stream_for_v2_verify,
renorm_draft_probs,
sample_draft_proposal,
select_top_k_tokens,
@@ -1468,214 +1460,21 @@ class EAGLEWorkerV2(BaseSpecWorker):
dw._rebuild_topk1_chain_buffers()
def verify(self, batch: ScheduleBatch):
fwd_stream = torch.get_device_module(self.device).current_stream()
verify_input: EagleVerifyInput = batch.spec_info
record_stream_for_v2_verify(batch, verify_input, fwd_stream)
bs = len(batch.seq_lens)
# Batch 1: Target verify
# Prepare for target verify in a separate stream
with self.plan_stream_ctx:
verify_forward_batch, can_run_cuda_graph = eagle_prepare_for_verify(
verify_input,
self.req_to_token_pool,
batch,
self.target_worker,
)
# Cover post-prepare rebinds: draft_token, plan_stream-allocated out_cache_loc.
record_stream_each((batch.input_ids, batch.out_cache_loc), fwd_stream)
# Correct some buffers due to the overlap plan
if self.plan_stream:
torch.get_device_module(self.device).current_stream().wait_stream(
self.plan_stream
)
if (
_is_npu
and self._target_worker.model_runner.model_config.model_is_mrope
and batch.spec_info is not None
and getattr(batch.spec_info, "positions", None) is not None
and not batch.forward_mode.is_idle()
):
# mrope_position depends on draft output in default stream and is computed in plan stream,
# causing errors. Compute it here for correct values.
verify_forward_batch.compute_spec_mrope_positions(
self._target_worker.model_runner, batch
)
# Some values such as custom_mask and position depend on the output of draft,
# so the previous plan step used the wrong values. Here, we need to run the related
# computation again to update them to the correct values.
self.target_worker.model_runner.attn_backend.update_verify_buffers_to_fill_after_draft(
verify_input,
(
self.target_worker.model_runner.decode_cuda_graph_runner.bs
if can_run_cuda_graph
else None
),
)
# Prepare grammar data on CPU if needed
if batch.has_grammar:
retrieve_next_token_cpu = verify_input.retrieve_next_token.cpu()
retrieve_next_sibling_cpu = verify_input.retrieve_next_sibling.cpu()
draft_tokens_cpu = verify_input.draft_token.view(
verify_input.retrieve_next_token.shape
).cpu()
# Run target verify batch in the main compute stream (GPU compute).
# Metadata init is skipped iff cuda-graph already ran load_batch —
# eagle_prepare_for_verify marked the batch in exactly that case; the
# non-cuda-graph path stays unmarked and gets forward_extend's init
# (post-pad).
forward_batch_output = self.target_worker.forward_batch_generation(
batch=None,
forward_batch=verify_forward_batch,
is_verify=True,
)
logits_output = forward_batch_output.logits_output
# Generate vocab mask for constrained decoding
vocab_mask = None
if batch.has_grammar:
# Generate the logit mask for structured output.
vocab_mask = generate_token_bitmask(
batch.reqs,
verify_input,
retrieve_next_token_cpu,
retrieve_next_sibling_cpu,
draft_tokens_cpu,
batch.sampling_info.vocab_size,
)
if vocab_mask is not None:
assert verify_input.grammar is not None
vocab_mask = vocab_mask.to(verify_input.retrieve_next_token.device)
# NOTE: otherwise, this vocab mask will be the one from the previous extend stage
# and will be applied to produce wrong results
batch.sampling_info.vocab_mask = None
# Sample
maybe_detect_nan(logits_output.next_token_logits, "verify: target model logits")
maybe_detect_inf(logits_output.next_token_logits, "verify: target model logits")
(
predict,
accept_lens,
accept_index,
) = eagle_sample(verify_input, batch, logits_output, vocab_mask)
new_seq_lens = batch.seq_lens + accept_lens
clear_unaccepted_c128 = getattr(
self.token_to_kv_pool_allocator.get_kvcache(),
"clear_unaccepted_c128_draft_states",
None,
)
if clear_unaccepted_c128 is not None and not batch.forward_mode.is_idle():
clear_unaccepted_c128(
batch.req_pool_indices,
batch.seq_lens,
accept_lens,
self.speculative_num_draft_tokens,
)
# Update mamba state for hybrid GDN models after verification
commit_mamba_states_after_verify(
self.target_worker,
return run_eagle_verify(
batch,
accept_lens,
accept_index,
self.speculative_num_draft_tokens,
target_worker=self.target_worker,
req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
plan_stream=self.plan_stream,
plan_stream_ctx=self.plan_stream_ctx,
topk=self.topk,
num_steps=self.speculative_num_steps,
num_draft_tokens=self.speculative_num_draft_tokens,
device=self.device,
metadata_ready_pre_pad=False,
finalize_tree_path=True,
)
if not batch.forward_mode.is_idle():
accept_tokens = predict[accept_index]
bonus_tokens = torch.empty_like(accept_lens, dtype=torch.int32)
# stride = accept_tokens per-req width = accept_index.shape[1]
# (spec_steps + 1); NOT num_draft_tokens, wrong for topk > 1 trees.
fill_bonus_tokens_func(
accept_tokens,
accept_lens,
bonus_tokens,
accept_index.shape[1],
bs,
)
else:
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(
batch, logits_output, predict, accept_index, self.speculative_num_steps
)
if not batch.forward_mode.is_idle() and self.topk > 1:
# topk == 1 needs nothing here: the accepted path is already the front
# chain, so the whole compaction is an identity transform.
predict = self._finalize_accept_tree_path(
batch, accept_index, accept_lens, predict, logits_output, bs
)
next_draft_input = EagleDraftInput(bonus_tokens=bonus_tokens)
# verify_forward_batch transitively holds verify-time GPU tensors
# (draft_token / out_cache_loc / ...) that must outlive the imminent
# batch.input_ids rebind in prepare_for_draft_extend.
# Scheduler pins it in batch_record_buf for the 2-iter window.
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=predict,
can_run_cuda_graph=can_run_cuda_graph,
speculative_num_draft_tokens=self.speculative_num_draft_tokens,
next_draft_input=next_draft_input,
accept_lens=accept_lens,
new_seq_lens=new_seq_lens,
routed_experts_output=forward_batch_output.routed_experts_output,
indexer_topk_output=forward_batch_output.indexer_topk_output,
extra_keep_alive_refs=[verify_forward_batch],
)
def _finalize_accept_tree_path(
self,
batch: ScheduleBatch,
accept_index: torch.Tensor,
accept_lens: torch.Tensor,
predict: torch.Tensor,
logits_output,
bs: int,
) -> torch.Tensor:
"""Tree drafting (topk > 1): move the accepted path -- KV slots, predict,
hidden_states -- to the contiguous front of each per-req block, which the
downstream chain-layout code (draft-extend select_index, committed-KV reads)
assumes. Returns compacted predict; mutates logits_output.hidden_states
(moved only when present)."""
move_accept_tokens_to_target_kvcache(
batch, accept_index, accept_lens - 1, self.token_to_kv_pool_allocator
)
predict = self._compact_accept_to_front(predict, accept_index, bs)
if logits_output.hidden_states is not None:
logits_output.hidden_states = self._compact_accept_to_front(
logits_output.hidden_states, accept_index, bs
)
return predict
def _compact_accept_to_front(
self, x: torch.Tensor, accept_index: torch.Tensor, bs: int
) -> torch.Tensor:
"""Gather the accepted tree path to the front of each per-req block.
``x`` is node-indexed over the whole tree (``[bs * num_draft_tokens, ...]``),
``accept_index`` is ``[bs, spec_steps + 1]`` global node indices (-1 padded).
Padded entries clamp to node 0 but land past accept_lens (never read);
trailing unaccepted slots stay and are freed as overshoot.
"""
nd = self.speculative_num_draft_tokens
s1 = accept_index.shape[1] # spec_steps + 1
safe = accept_index.to(torch.int64).clamp(min=0).reshape(-1)
gathered = x[safe]
out = x.clone()
out.view(bs, nd, *x.shape[1:])[:, :s1] = gathered.view(bs, s1, *x.shape[1:])
return out
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
monkey_patch_torch_reductions()
named_tensors = MultiprocessingSerializer.deserialize(
@@ -20,14 +20,12 @@ from typing import TYPE_CHECKING, List
import torch
from sglang.kernels.ops.speculative.eagle import fill_bonus_tokens_func
from sglang.srt.distributed.parallel_state_wrapper import ParallelState
from sglang.srt.environ import envs
from sglang.srt.hardware_backend.npu.graph_runner.multi_layer_eagle_draft_extend_npu_graph_runner import (
MultiLayerEagleMultiStepDraftExtendNpuGraphRunner,
)
from sglang.srt.layers.moe.utils import speculative_moe_backend_context
from sglang.srt.layers.utils.logprob import compute_spec_v2_logprobs
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.scheduler import GenerationBatchResult
from sglang.srt.managers.tp_worker import TpModelWorker
@@ -50,14 +48,13 @@ from sglang.srt.speculative.eagle_info import (
)
from sglang.srt.speculative.eagle_utils import (
default_tree_mask_mode,
eagle_prepare_for_verify,
eagle_sample,
get_draft_recurrent_hidden_state_spec,
)
from sglang.srt.speculative.eagle_worker_common import (
build_eagle_verify_input,
prepare_for_draft,
prepare_for_draft_extend,
run_eagle_verify,
)
from sglang.srt.speculative.multi_layer_eagle_draft_extend_cuda_graph_runner import (
MultiLayerEagleMultiStepDraftExtendCudaGraphRunner,
@@ -67,8 +64,6 @@ from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
from sglang.srt.speculative.spec_utils import (
draft_tp_context,
get_plan_stream,
record_stream_each,
record_stream_for_v2_verify,
sample_draft_proposal,
select_top_k_tokens,
)
@@ -723,104 +718,18 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
self.draft_worker._draft_extend_for_decode(batch, batch_output)
return batch_output
def verify(
self,
batch: ScheduleBatch,
):
fwd_stream = torch.get_device_module(self.device).current_stream()
verify_input: EagleVerifyInput = batch.spec_info
record_stream_for_v2_verify(batch, verify_input, fwd_stream)
bs = len(batch.seq_lens)
# Batch 1: Target verify
# Prepare for target verify in a separate stream
with self.plan_stream_ctx:
verify_forward_batch, can_run_cuda_graph = eagle_prepare_for_verify(
verify_input,
self.req_to_token_pool,
batch,
self.target_worker,
)
# Cover post-prepare rebinds: draft_token, plan_stream-allocated out_cache_loc.
record_stream_each((batch.input_ids, batch.out_cache_loc), fwd_stream)
# Correct some buffers due to the overlap plan
if self.plan_stream:
torch.get_device_module(self.device).current_stream().wait_stream(
self.plan_stream
)
# Some values such as custom_mask and position depend on the output of draft,
# so the previous plan step used the wrong values. Here, we need to run the related
# computation again to update them to the correct values.
self.target_worker.model_runner.attn_backend.update_verify_buffers_to_fill_after_draft(
verify_input,
(
self.target_worker.model_runner.decode_cuda_graph_runner.bs
if can_run_cuda_graph
else None
),
)
# NOTE: metadata init is skipped here unconditionally, although
# eagle_prepare_for_verify only plans when cuda-graph load_batch ran.
# eagle_worker_v2 re-inits the non-graph path instead (post-pad); this
# worker has not adopted that fix, so preserve its behavior verbatim.
# On NPU with --disable-cuda-graph, non-graph verify needs metadata init
# in forward_extend (post-pad); only mark ready for the cuda-graph path.
if not _is_npu or can_run_cuda_graph:
verify_forward_batch.mark_forward_metadata_ready()
# Run target verify batch in the main compute stream
forward_batch_output = self.target_worker.forward_batch_generation(
batch=None,
forward_batch=verify_forward_batch,
is_verify=True,
)
logits_output = forward_batch_output.logits_output
# Sample
maybe_detect_nan(logits_output.next_token_logits, "verify: target model logits")
maybe_detect_inf(logits_output.next_token_logits, "verify: target model logits")
(
predict,
accept_lens,
accept_index,
) = eagle_sample(verify_input, batch, logits_output)
new_seq_lens = batch.seq_lens + accept_lens
if not batch.forward_mode.is_idle():
accept_tokens = predict[accept_index]
bonus_tokens = torch.empty_like(accept_lens, dtype=torch.int32)
# stride = accept_tokens per-req width = accept_index.shape[1].
fill_bonus_tokens_func(
accept_tokens,
accept_lens,
bonus_tokens,
accept_index.shape[1],
bs,
)
else:
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(
batch, logits_output, predict, accept_index, self.speculative_num_steps
)
next_draft_input = EagleDraftInput(bonus_tokens=bonus_tokens)
# verify_forward_batch transitively holds verify-time GPU tensors that
# must outlive the imminent batch.input_ids rebind; scheduler pins it
# in batch_record_buf via extra_keep_alive_refs. See EAGLEWorkerV2.verify.
return GenerationBatchResult(
logits_output=logits_output,
next_token_ids=predict,
can_run_cuda_graph=can_run_cuda_graph,
speculative_num_draft_tokens=self.speculative_num_draft_tokens,
next_draft_input=next_draft_input,
accept_lens=accept_lens,
new_seq_lens=new_seq_lens,
routed_experts_output=forward_batch_output.routed_experts_output,
indexer_topk_output=forward_batch_output.indexer_topk_output,
extra_keep_alive_refs=[verify_forward_batch],
def verify(self, batch: ScheduleBatch):
return run_eagle_verify(
batch,
target_worker=self.target_worker,
req_to_token_pool=self.req_to_token_pool,
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
plan_stream=self.plan_stream,
plan_stream_ctx=self.plan_stream_ctx,
topk=self.topk,
num_steps=self.speculative_num_steps,
num_draft_tokens=self.speculative_num_draft_tokens,
device=self.device,
metadata_ready_pre_pad=True,
finalize_tree_path=False,
)