[Spec] Consolidate the verify step into eagle_worker_common.run_eagle_verify (#31380)
This commit is contained in:
@@ -7,21 +7,38 @@ import torch
|
|||||||
from sglang.kernels.ops.speculative.cache_locs import (
|
from sglang.kernels.ops.speculative.cache_locs import (
|
||||||
assign_draft_cache_locs_contiguous,
|
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 (
|
from sglang.srt.model_executor.forward_batch_info import (
|
||||||
CaptureHiddenMode,
|
CaptureHiddenMode,
|
||||||
ForwardBatch,
|
ForwardBatch,
|
||||||
ForwardMode,
|
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 (
|
from sglang.srt.speculative.eagle_utils import (
|
||||||
TreeMaskMode,
|
TreeMaskMode,
|
||||||
build_tree_kernel_efficient,
|
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 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
|
from sglang.srt.utils.common import is_npu
|
||||||
|
|
||||||
_is_cpu = is_cpu()
|
_is_cpu = is_cpu()
|
||||||
|
_is_npu = is_npu()
|
||||||
|
|
||||||
if _is_cpu:
|
if _is_cpu:
|
||||||
from sgl_kernel import assign_draft_cache_locs_contiguous_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 (
|
from sglang.srt.speculative.eagle_draft_cuda_graph_runner import (
|
||||||
EAGLEDraftCudaGraphRunner,
|
EAGLEDraftCudaGraphRunner,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.eagle_info import (
|
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
|
||||||
EagleDraftExtendInput,
|
|
||||||
EagleDraftInput,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
def duplicate_prefix_tail_to_draft_branches(
|
def duplicate_prefix_tail_to_draft_branches(
|
||||||
@@ -357,3 +371,266 @@ def build_eagle_verify_input(
|
|||||||
seq_lens_cpu=None,
|
seq_lens_cpu=None,
|
||||||
draft_probs=draft_probs,
|
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],
|
||||||
|
)
|
||||||
|
|||||||
@@ -6,7 +6,6 @@ from typing import List, Optional
|
|||||||
|
|
||||||
import torch
|
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.distributed.parallel_state_wrapper import ParallelState
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.hardware_backend.npu.graph_runner.eagle_draft_extend_npu_graph_runner import (
|
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_a2a_backend_context,
|
||||||
speculative_moe_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.io_struct import UpdateWeightsFromTensorReqInput
|
||||||
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
from sglang.srt.managers.schedule_batch import ScheduleBatch
|
||||||
from sglang.srt.managers.scheduler import GenerationBatchResult
|
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 (
|
from sglang.srt.speculative.eagle_utils import (
|
||||||
_eagle_prefill_tail_tokens,
|
_eagle_prefill_tail_tokens,
|
||||||
default_tree_mask_mode,
|
default_tree_mask_mode,
|
||||||
eagle_prepare_for_verify,
|
|
||||||
eagle_sample,
|
|
||||||
get_draft_recurrent_hidden_state_spec,
|
get_draft_recurrent_hidden_state_spec,
|
||||||
organize_draft_results,
|
organize_draft_results,
|
||||||
per_step_draft_out_cache_loc,
|
per_step_draft_out_cache_loc,
|
||||||
@@ -76,18 +72,14 @@ from sglang.srt.speculative.eagle_worker_common import (
|
|||||||
build_eagle_verify_input,
|
build_eagle_verify_input,
|
||||||
prepare_for_draft,
|
prepare_for_draft,
|
||||||
prepare_for_draft_extend,
|
prepare_for_draft_extend,
|
||||||
|
run_eagle_verify,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
||||||
from sglang.srt.speculative.spec_utils import (
|
from sglang.srt.speculative.spec_utils import (
|
||||||
commit_mamba_states_after_verify,
|
|
||||||
draft_tp_context,
|
draft_tp_context,
|
||||||
fast_sample,
|
fast_sample,
|
||||||
generate_token_bitmask,
|
|
||||||
get_plan_stream,
|
get_plan_stream,
|
||||||
load_token_map,
|
load_token_map,
|
||||||
move_accept_tokens_to_target_kvcache,
|
|
||||||
record_stream_each,
|
|
||||||
record_stream_for_v2_verify,
|
|
||||||
renorm_draft_probs,
|
renorm_draft_probs,
|
||||||
sample_draft_proposal,
|
sample_draft_proposal,
|
||||||
select_top_k_tokens,
|
select_top_k_tokens,
|
||||||
@@ -1468,214 +1460,21 @@ class EAGLEWorkerV2(BaseSpecWorker):
|
|||||||
dw._rebuild_topk1_chain_buffers()
|
dw._rebuild_topk1_chain_buffers()
|
||||||
|
|
||||||
def verify(self, batch: ScheduleBatch):
|
def verify(self, batch: ScheduleBatch):
|
||||||
fwd_stream = torch.get_device_module(self.device).current_stream()
|
return run_eagle_verify(
|
||||||
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,
|
|
||||||
batch,
|
batch,
|
||||||
accept_lens,
|
target_worker=self.target_worker,
|
||||||
accept_index,
|
req_to_token_pool=self.req_to_token_pool,
|
||||||
self.speculative_num_draft_tokens,
|
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):
|
def update_weights_from_tensor(self, recv_req: UpdateWeightsFromTensorReqInput):
|
||||||
monkey_patch_torch_reductions()
|
monkey_patch_torch_reductions()
|
||||||
named_tensors = MultiprocessingSerializer.deserialize(
|
named_tensors = MultiprocessingSerializer.deserialize(
|
||||||
|
|||||||
@@ -20,14 +20,12 @@ from typing import TYPE_CHECKING, List
|
|||||||
|
|
||||||
import torch
|
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.distributed.parallel_state_wrapper import ParallelState
|
||||||
from sglang.srt.environ import envs
|
from sglang.srt.environ import envs
|
||||||
from sglang.srt.hardware_backend.npu.graph_runner.multi_layer_eagle_draft_extend_npu_graph_runner import (
|
from sglang.srt.hardware_backend.npu.graph_runner.multi_layer_eagle_draft_extend_npu_graph_runner import (
|
||||||
MultiLayerEagleMultiStepDraftExtendNpuGraphRunner,
|
MultiLayerEagleMultiStepDraftExtendNpuGraphRunner,
|
||||||
)
|
)
|
||||||
from sglang.srt.layers.moe.utils import speculative_moe_backend_context
|
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.schedule_batch import ScheduleBatch
|
||||||
from sglang.srt.managers.scheduler import GenerationBatchResult
|
from sglang.srt.managers.scheduler import GenerationBatchResult
|
||||||
from sglang.srt.managers.tp_worker import TpModelWorker
|
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 (
|
from sglang.srt.speculative.eagle_utils import (
|
||||||
default_tree_mask_mode,
|
default_tree_mask_mode,
|
||||||
eagle_prepare_for_verify,
|
|
||||||
eagle_sample,
|
|
||||||
get_draft_recurrent_hidden_state_spec,
|
get_draft_recurrent_hidden_state_spec,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.eagle_worker_common import (
|
from sglang.srt.speculative.eagle_worker_common import (
|
||||||
build_eagle_verify_input,
|
build_eagle_verify_input,
|
||||||
prepare_for_draft,
|
prepare_for_draft,
|
||||||
prepare_for_draft_extend,
|
prepare_for_draft_extend,
|
||||||
|
run_eagle_verify,
|
||||||
)
|
)
|
||||||
from sglang.srt.speculative.multi_layer_eagle_draft_extend_cuda_graph_runner import (
|
from sglang.srt.speculative.multi_layer_eagle_draft_extend_cuda_graph_runner import (
|
||||||
MultiLayerEagleMultiStepDraftExtendCudaGraphRunner,
|
MultiLayerEagleMultiStepDraftExtendCudaGraphRunner,
|
||||||
@@ -67,8 +64,6 @@ from sglang.srt.speculative.spec_info import SpeculativeAlgorithm
|
|||||||
from sglang.srt.speculative.spec_utils import (
|
from sglang.srt.speculative.spec_utils import (
|
||||||
draft_tp_context,
|
draft_tp_context,
|
||||||
get_plan_stream,
|
get_plan_stream,
|
||||||
record_stream_each,
|
|
||||||
record_stream_for_v2_verify,
|
|
||||||
sample_draft_proposal,
|
sample_draft_proposal,
|
||||||
select_top_k_tokens,
|
select_top_k_tokens,
|
||||||
)
|
)
|
||||||
@@ -723,104 +718,18 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
|
|||||||
self.draft_worker._draft_extend_for_decode(batch, batch_output)
|
self.draft_worker._draft_extend_for_decode(batch, batch_output)
|
||||||
return batch_output
|
return batch_output
|
||||||
|
|
||||||
def verify(
|
def verify(self, batch: ScheduleBatch):
|
||||||
self,
|
return run_eagle_verify(
|
||||||
batch: ScheduleBatch,
|
batch,
|
||||||
):
|
target_worker=self.target_worker,
|
||||||
fwd_stream = torch.get_device_module(self.device).current_stream()
|
req_to_token_pool=self.req_to_token_pool,
|
||||||
verify_input: EagleVerifyInput = batch.spec_info
|
token_to_kv_pool_allocator=self.token_to_kv_pool_allocator,
|
||||||
record_stream_for_v2_verify(batch, verify_input, fwd_stream)
|
plan_stream=self.plan_stream,
|
||||||
|
plan_stream_ctx=self.plan_stream_ctx,
|
||||||
bs = len(batch.seq_lens)
|
topk=self.topk,
|
||||||
|
num_steps=self.speculative_num_steps,
|
||||||
# Batch 1: Target verify
|
num_draft_tokens=self.speculative_num_draft_tokens,
|
||||||
# Prepare for target verify in a separate stream
|
device=self.device,
|
||||||
with self.plan_stream_ctx:
|
metadata_ready_pre_pad=True,
|
||||||
verify_forward_batch, can_run_cuda_graph = eagle_prepare_for_verify(
|
finalize_tree_path=False,
|
||||||
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],
|
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user