[Spec] Extract the shared draft() tail into build_eagle_verify_input (#31375)

This commit is contained in:
Liangsheng Yin
2026-07-15 15:59:31 -07:00
committed by GitHub
parent 7a973c03a0
commit b0b2dfbda1
4 changed files with 114 additions and 124 deletions
@@ -1,6 +1,6 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Any
from typing import TYPE_CHECKING, Any, Optional
import torch
@@ -12,6 +12,11 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch,
ForwardMode,
)
from sglang.srt.speculative.eagle_info import EagleVerifyInput
from sglang.srt.speculative.eagle_utils import (
TreeMaskMode,
build_tree_kernel_efficient,
)
from sglang.srt.utils import is_cpu
from sglang.srt.utils.async_probe import maybe_detect_oob
from sglang.srt.utils.common import is_npu
@@ -23,6 +28,7 @@ if _is_cpu:
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.mem_cache.memory_pool import ReqToTokenPool
from sglang.srt.model_executor.model_runner import ModelRunner
from sglang.srt.speculative.eagle_draft_cuda_graph_runner import (
@@ -265,3 +271,89 @@ def prepare_for_draft(
forward_batch
)
return forward_batch, can_cuda_graph
def build_eagle_verify_input(
batch: ScheduleBatch,
draft_input: EagleDraftInput,
parent_list: torch.Tensor,
top_scores_index: torch.Tensor,
draft_tokens: torch.Tensor,
draft_probs: Optional[torch.Tensor],
*,
target_worker: TpModelWorker,
topk: int,
num_steps: int,
num_draft_tokens: int,
tree_mask_mode: TreeMaskMode,
device: str,
) -> EagleVerifyInput:
"""Shared draft() tail: idle input, tree-mask build, EagleVerifyInput assembly.
``draft_probs`` is the caller's source of draft distributions (single-layer
eagle: this round's draft_forward output; multi-layer eagle: the ones the
draft input carried).
"""
if batch.forward_mode.is_idle():
return EagleVerifyInput.create_idle_input(
topk,
num_steps,
num_draft_tokens,
device,
)
# Build tree mask
# Directly write to cuda graph buffers for verify attn
tree_mask_buf, position_buf = (
target_worker.model_runner.attn_backend.get_verify_buffers_to_fill_after_draft()
)
# build_tree_kernel uses seq_lens_sum only to size the (non-preallocated)
# tree mask; over-size is safe. Skip per-iter .sum().item() D2H via UB.
seq_lens_sum = batch.seq_lens_sum
if seq_lens_sum is None:
if tree_mask_buf is None:
max_context_len = target_worker.model_runner.attn_backend.max_context_len
seq_lens_sum = batch.seq_lens.shape[0] * max_context_len
else:
# tree_mask_buf preallocated -> kernel ignores seq_lens_sum.
seq_lens_sum = 0
(
tree_mask,
position,
retrieve_index,
retrieve_next_token,
retrieve_next_sibling,
draft_tokens,
) = build_tree_kernel_efficient(
draft_input.bonus_tokens,
parent_list,
top_scores_index,
draft_tokens,
batch.seq_lens,
seq_lens_sum,
topk,
num_steps,
num_draft_tokens,
tree_mask_mode,
tree_mask_buf,
position_buf,
)
return EagleVerifyInput(
draft_token=draft_tokens,
custom_mask=tree_mask,
positions=position,
retrieve_index=retrieve_index,
retrieve_next_token=retrieve_next_token,
retrieve_next_sibling=retrieve_next_sibling,
retrieve_cum_len=None,
spec_steps=num_steps,
topk=topk,
draft_token_num=num_draft_tokens,
capture_hidden_mode=None,
seq_lens_sum=None,
seq_lens_cpu=None,
draft_probs=draft_probs,
)
@@ -65,7 +65,6 @@ from sglang.srt.speculative.eagle_info import (
)
from sglang.srt.speculative.eagle_utils import (
_eagle_prefill_tail_tokens,
build_tree_kernel_efficient,
default_tree_mask_mode,
eagle_prepare_for_verify,
eagle_sample,
@@ -74,6 +73,7 @@ from sglang.srt.speculative.eagle_utils import (
per_step_draft_out_cache_loc,
)
from sglang.srt.speculative.eagle_worker_common import (
build_eagle_verify_input,
prepare_for_draft,
prepare_for_draft_extend,
)
@@ -546,70 +546,19 @@ class EagleDraftWorker(EagleDraftWorkerBase):
self.draft_forward(forward_batch)
)
if batch.forward_mode.is_idle():
return EagleVerifyInput.create_idle_input(
self.topk,
self.speculative_num_steps,
self.speculative_num_draft_tokens,
self.device,
)
# Build tree mask
# Directly write to cuda graph buffers for verify attn
tree_mask_buf, position_buf = (
self.target_worker.model_runner.attn_backend.get_verify_buffers_to_fill_after_draft()
)
# build_tree_kernel uses seq_lens_sum only to size the (non-preallocated)
# tree mask; over-size is safe. Skip per-iter .sum().item() D2H via UB.
seq_lens_sum = batch.seq_lens_sum
if seq_lens_sum is None:
if tree_mask_buf is None:
max_context_len = (
self.target_worker.model_runner.attn_backend.max_context_len
)
seq_lens_sum = batch.seq_lens.shape[0] * max_context_len
else:
# tree_mask_buf preallocated -> kernel ignores seq_lens_sum.
seq_lens_sum = 0
(
tree_mask,
position,
retrieve_index,
retrieve_next_token,
retrieve_next_sibling,
draft_tokens,
) = build_tree_kernel_efficient(
draft_input.bonus_tokens,
return build_eagle_verify_input(
batch,
draft_input,
parent_list,
top_scores_index,
draft_tokens,
batch.seq_lens,
seq_lens_sum,
self.topk,
self.speculative_num_steps,
self.speculative_num_draft_tokens,
self.tree_mask_mode,
tree_mask_buf,
position_buf,
)
return EagleVerifyInput(
draft_token=draft_tokens,
custom_mask=tree_mask,
positions=position,
retrieve_index=retrieve_index,
retrieve_next_token=retrieve_next_token,
retrieve_next_sibling=retrieve_next_sibling,
retrieve_cum_len=None,
spec_steps=self.speculative_num_steps,
draft_probs,
target_worker=self.target_worker,
topk=self.topk,
draft_token_num=self.speculative_num_draft_tokens,
capture_hidden_mode=None,
seq_lens_sum=None,
seq_lens_cpu=None,
draft_probs=draft_probs,
num_steps=self.speculative_num_steps,
num_draft_tokens=self.speculative_num_draft_tokens,
tree_mask_mode=self.tree_mask_mode,
device=self.device,
)
def draft_forward(self, forward_batch: ForwardBatch):
@@ -49,13 +49,13 @@ from sglang.srt.speculative.eagle_info import (
EagleVerifyInput,
)
from sglang.srt.speculative.eagle_utils import (
build_tree_kernel_efficient,
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,
)
@@ -252,70 +252,19 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
# Run draft
parent_list, top_scores_index, draft_tokens = self.draft_forward(forward_batch)
if batch.forward_mode.is_idle():
return EagleVerifyInput.create_idle_input(
self.topk,
self.speculative_num_steps,
self.speculative_num_draft_tokens,
self.device,
)
# Build tree mask
# Directly write to cuda graph buffers for verify attn
tree_mask_buf, position_buf = (
self.target_worker.model_runner.attn_backend.get_verify_buffers_to_fill_after_draft()
)
# build_tree_kernel uses seq_lens_sum only to size the (non-preallocated)
# tree mask; over-size is safe. Skip per-iter .sum().item() D2H via UB.
seq_lens_sum = batch.seq_lens_sum
if seq_lens_sum is None:
if tree_mask_buf is None:
max_context_len = (
self.target_worker.model_runner.attn_backend.max_context_len
)
seq_lens_sum = batch.seq_lens.shape[0] * max_context_len
else:
# tree_mask_buf preallocated -> kernel ignores seq_lens_sum.
seq_lens_sum = 0
(
tree_mask,
position,
retrieve_index,
retrieve_next_token,
retrieve_next_sibling,
draft_tokens,
) = build_tree_kernel_efficient(
draft_input.bonus_tokens,
return build_eagle_verify_input(
batch,
draft_input,
parent_list,
top_scores_index,
draft_tokens,
batch.seq_lens,
seq_lens_sum,
self.topk,
self.speculative_num_steps,
self.speculative_num_draft_tokens,
self.tree_mask_mode,
tree_mask_buf,
position_buf,
)
return EagleVerifyInput(
draft_token=draft_tokens,
custom_mask=tree_mask,
positions=position,
retrieve_index=retrieve_index,
retrieve_next_token=retrieve_next_token,
retrieve_next_sibling=retrieve_next_sibling,
retrieve_cum_len=None,
spec_steps=self.speculative_num_steps,
draft_input.draft_probs,
target_worker=self.target_worker,
topk=self.topk,
draft_token_num=self.speculative_num_draft_tokens,
capture_hidden_mode=None,
seq_lens_sum=None,
seq_lens_cpu=None,
draft_probs=draft_input.draft_probs,
num_steps=self.speculative_num_steps,
num_draft_tokens=self.speculative_num_draft_tokens,
tree_mask_mode=self.tree_mask_mode,
device=self.device,
)
def draft_forward(self, forward_batch: ForwardBatch):
@@ -192,7 +192,7 @@ class TestEagleWorkerV2BackendFallback(CustomTestCase):
)
with patch(
"sglang.srt.speculative.eagle_worker_v2.build_tree_kernel_efficient",
"sglang.srt.speculative.eagle_worker_common.build_tree_kernel_efficient",
return_value=tree_result,
), patch(
"sglang.srt.speculative.eagle_worker_v2.prepare_for_draft",