[Spec] Move draft-extend prep to EagleDraftWorkerBase; unify prepare_for_* names (#28093)

This commit is contained in:
Liangsheng Yin
2026-06-12 16:44:04 -07:00
committed by GitHub
parent f23f48df98
commit 87554c7855
12 changed files with 126 additions and 119 deletions
+2 -2
View File
@@ -111,14 +111,14 @@ def _extract_prefix_lens_and_extend_seq_lens(
out_prefix_lens[: positions.shape[0]].copy_(positions.to(torch.int64))
out_extend_seq_lens.fill_(1)
elif forward_mode.is_target_verify():
# Evidence: EagleVerifyInputV2Mixin.prepare_for_v2_verify assigns out_cache_loc in
# Evidence: EagleVerifyInputV2Mixin.prepare_for_verify assigns out_cache_loc in
# [seq_lens, seq_lens + draft_token_num) without bumping seq_lens. The target-verify
# branch in TRTLLMHAAttnBackend.init_forward_metadata uses seq_lens as the prefix and
# tokens_per_req as the query length, so mirror that as seq_lens plus draft_token_num.
out_prefix_lens.copy_(forward_batch.seq_lens[:bs].to(torch.int64))
out_extend_seq_lens.fill_(int(spec_info.draft_token_num))
elif forward_mode.is_draft_extend_v2():
# Evidence: EagleDraftExtendInputV2Mixin.prepare_for_extend_to_fill_draft_kvcache bumps
# Evidence: EagleDraftWorkerBase.prepare_for_draft_extend bumps
# seq_lens by num_draft_tokens. FlashAttentionBackend.init_forward_metadata reads the
# draft-extend-v2 query length from spec_info.extend_seq_lens_tensor when available.
# CUDA-graph replay passes extend_seq_lens but omits extend_prefix_lens, so derive the
@@ -1,13 +1,17 @@
from __future__ import annotations
from abc import ABC, abstractmethod
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any
import torch
if TYPE_CHECKING:
from sglang.srt.managers.schedule_batch import ScheduleBatch
from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.speculative.eagle_info import EagleDraftExtendInput
class BaseDraftWorker(ABC):
class EagleDraftWorkerBase(ABC):
@abstractmethod
def draft():
pass
@@ -29,6 +33,84 @@ class BaseDraftWorker(ABC):
self.init_attention_backend()
self.init_cuda_graphs()
def prepare_for_draft_extend(
self,
draft_extend_input: EagleDraftExtendInput,
batch: ScheduleBatch,
predict: torch.Tensor,
num_draft_tokens: int,
draft_model_runner: Any,
cuda_graph_runner: Any,
):
from sglang.srt.model_executor.forward_batch_info import (
CaptureHiddenMode,
ForwardBatch,
ForwardMode,
)
from sglang.srt.utils.async_probe import maybe_detect_oob
from sglang.srt.utils.common import is_npu
bs = len(batch.seq_lens)
extend_num_tokens = bs * num_draft_tokens
# When seq_lens_cpu is absent, stay on GPU-only path -- no .tolist()/.cpu().
gpu_only = batch.seq_lens_cpu is None
batch.spec_info = draft_extend_input
batch.input_ids = predict
maybe_detect_oob(
batch.input_ids,
0,
batch.model_config.vocab_size,
"v2 prepare_for_draft_extend input_ids",
)
# init_new requires both list or both Tensor;
# gpu_only emits device tensors to skip H2D.
if gpu_only:
batch.prefix_lens = batch.seq_lens.to(torch.int32)
batch.extend_lens = torch.full(
(bs,), num_draft_tokens, dtype=torch.int32, device=batch.seq_lens.device
)
else:
batch.prefix_lens = batch.seq_lens_cpu.tolist()
batch.extend_lens = [num_draft_tokens] * bs
batch.extend_num_tokens = extend_num_tokens
capture_mode = (
CaptureHiddenMode.NULL
if draft_model_runner.spec_algorithm.is_standalone()
else CaptureHiddenMode.FULL
)
batch.forward_mode = (
ForwardMode.IDLE
if batch.forward_mode.is_idle()
else ForwardMode.DRAFT_EXTEND_V2
)
batch.capture_hidden_mode = capture_mode
forward_batch = ForwardBatch.init_new(batch, draft_model_runner)
# Forward sees post-write length (draft extend writes num_draft_tokens
# slots); mutation stays on forward_batch to preserve SB.seq_lens.
forward_batch.seq_lens = forward_batch.seq_lens + num_draft_tokens
if not gpu_only:
forward_batch.seq_lens_cpu = forward_batch.seq_lens_cpu + num_draft_tokens
forward_batch.seq_lens_sum = int(forward_batch.seq_lens_cpu.sum())
else:
# Supply CPU mirror (extend_seq_lens are all num_draft_tokens) so
# backend max() reads from list without a per-iter D2H sync.
forward_batch.extend_seq_lens_cpu = [num_draft_tokens] * bs
can_cuda_graph = cuda_graph_runner and cuda_graph_runner.can_run(forward_batch)
if not batch.forward_mode.is_idle() and not can_cuda_graph:
draft_model_runner.attn_backend.init_forward_metadata(forward_batch)
# Planned pre-pad; do NOT opt into post-pad re-plan. DSA's indexer
# cannot rebuild its deep_gemm schedule_meta on a DP-padded batch
# (the `_batch_size == batch_size` assertion, see #27091); the
# marked pre-pad metadata is used as-is, matching the proven
# skip_attn_backend_init=True behavior.
# On NPU with --disable-cuda-graph, block_table shape won't match
# after prepare_mlp_sync_batch padding; defer re-init to
# forward_extend (post-pad) instead.
if not is_npu() or can_cuda_graph:
forward_batch.mark_forward_metadata_ready()
return forward_batch
class BaseSpecWorker(ABC):
@property
@@ -38,7 +120,7 @@ class BaseSpecWorker(ABC):
@property
@abstractmethod
def draft_worker(self) -> BaseDraftWorker:
def draft_worker(self) -> EagleDraftWorkerBase:
pass
@property
+1 -1
View File
@@ -49,7 +49,7 @@ class DFlashVerifyInput(SpecInput):
def get_spec_adjust_token_coefficient(self) -> Tuple[int, int]:
return self.draft_token_num, self.draft_token_num
def prepare_for_v2_verify(
def prepare_for_verify(
self,
batch: ScheduleBatch,
target_worker: TpModelWorker,
@@ -262,7 +262,7 @@ class DFlashWorkerV2(BaseSpecWorker):
def draft_worker(self):
# DFLASH drives the draft model through a plain TpModelWorker: the
# draft KV is materialized from target hidden states, so there is no
# BaseDraftWorker draft/draft_extend split to wrap it in.
# EagleDraftWorkerBase draft/draft_extend split to wrap it in.
return self._draft_worker
@property
@@ -1524,7 +1524,7 @@ class DFlashWorkerV2(BaseSpecWorker):
model_worker_batch.seq_lens_cpu = draft_input.reserved_seq_lens_cpu
model_worker_batch.seq_lens_sum = int(draft_input.reserved_seq_lens_sum)
verify_forward_batch, _ = verify_input.prepare_for_v2_verify(
verify_forward_batch, _ = verify_input.prepare_for_verify(
model_worker_batch, self.target_worker
)
model_worker_batch.seq_lens_cpu = seq_lens_cpu_backup
+1 -2
View File
@@ -9,7 +9,6 @@ from sglang.srt.environ import envs
from sglang.srt.layers.attention.utils import create_flashinfer_kv_indices_triton
from sglang.srt.model_executor.forward_batch_info import CaptureHiddenMode
from sglang.srt.speculative.eagle_info_v2 import (
EagleDraftExtendInputV2Mixin,
EagleDraftInputV2Mixin,
EagleVerifyInputV2Mixin,
)
@@ -273,7 +272,7 @@ class EagleDraftInput(SpecInput, EagleDraftInputV2Mixin):
@dataclass
class EagleDraftExtendInput(SpecInput, EagleDraftExtendInputV2Mixin):
class EagleDraftExtendInput(SpecInput):
"""Inputs to the draft-extend forward (the fill-draft-kvcache pass after
target prefill / verify).
+3 -75
View File
@@ -1,7 +1,7 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import TYPE_CHECKING, Any
from typing import TYPE_CHECKING
import torch
import torch.nn.functional as F
@@ -61,7 +61,6 @@ if TYPE_CHECKING:
EAGLEDraftCudaGraphRunner,
)
from sglang.srt.speculative.eagle_info import (
EagleDraftExtendInput,
EagleDraftInput,
EagleVerifyInput,
)
@@ -220,7 +219,7 @@ class EagleDraftInputV2Mixin:
bs,
)
def prepare_for_v2_draft(
def prepare_for_draft(
self: EagleDraftInput,
req_to_token_pool: ReqToTokenPool,
batch: ScheduleBatch,
@@ -307,77 +306,6 @@ class EagleDraftInputV2Mixin:
return forward_batch, can_cuda_graph
class EagleDraftExtendInputV2Mixin:
def prepare_for_extend_to_fill_draft_kvcache(
self: EagleDraftExtendInput,
batch: ScheduleBatch,
predict: torch.Tensor,
num_draft_tokens: int,
draft_model_runner: Any,
cuda_graph_runner: Any,
):
bs = len(batch.seq_lens)
extend_num_tokens = bs * num_draft_tokens
# When seq_lens_cpu is absent, stay on GPU-only path -- no .tolist()/.cpu().
gpu_only = batch.seq_lens_cpu is None
batch.spec_info = self
batch.input_ids = predict
maybe_detect_oob(
batch.input_ids,
0,
batch.model_config.vocab_size,
"v2 prepare_for_extend_to_fill_draft_kvcache input_ids",
)
# init_new requires both list or both Tensor;
# gpu_only emits device tensors to skip H2D.
if gpu_only:
batch.prefix_lens = batch.seq_lens.to(torch.int32)
batch.extend_lens = torch.full(
(bs,), num_draft_tokens, dtype=torch.int32, device=batch.seq_lens.device
)
else:
batch.prefix_lens = batch.seq_lens_cpu.tolist()
batch.extend_lens = [num_draft_tokens] * bs
batch.extend_num_tokens = extend_num_tokens
capture_mode = (
CaptureHiddenMode.NULL
if draft_model_runner.spec_algorithm.is_standalone()
else CaptureHiddenMode.FULL
)
batch.forward_mode = (
ForwardMode.IDLE
if batch.forward_mode.is_idle()
else ForwardMode.DRAFT_EXTEND_V2
)
batch.capture_hidden_mode = capture_mode
forward_batch = ForwardBatch.init_new(batch, draft_model_runner)
# Forward sees post-write length (draft extend writes num_draft_tokens
# slots); mutation stays on forward_batch to preserve SB.seq_lens.
forward_batch.seq_lens = forward_batch.seq_lens + num_draft_tokens
if not gpu_only:
forward_batch.seq_lens_cpu = forward_batch.seq_lens_cpu + num_draft_tokens
forward_batch.seq_lens_sum = int(forward_batch.seq_lens_cpu.sum())
else:
# Supply CPU mirror (extend_seq_lens are all num_draft_tokens) so
# backend max() reads from list without a per-iter D2H sync.
forward_batch.extend_seq_lens_cpu = [num_draft_tokens] * bs
can_cuda_graph = cuda_graph_runner and cuda_graph_runner.can_run(forward_batch)
if not batch.forward_mode.is_idle() and not can_cuda_graph:
draft_model_runner.attn_backend.init_forward_metadata(forward_batch)
# Planned pre-pad; do NOT opt into post-pad re-plan. DSA's indexer
# cannot rebuild its deep_gemm schedule_meta on a DP-padded batch
# (the `_batch_size == batch_size` assertion, see #27091); the
# marked pre-pad metadata is used as-is, matching the proven
# skip_attn_backend_init=True behavior.
# On NPU with --disable-cuda-graph, block_table shape won't match
# after prepare_mlp_sync_batch padding; defer re-init to
# forward_extend (post-pad) instead.
if not _is_npu or can_cuda_graph:
forward_batch.mark_forward_metadata_ready()
return forward_batch
@dataclass
class EagleVerifyInputV2Mixin:
@property
@@ -393,7 +321,7 @@ class EagleVerifyInputV2Mixin:
irregular tree (no fixed per-level branching)."""
return self.topk
def prepare_for_v2_verify(
def prepare_for_verify(
self: EagleVerifyInput,
req_to_token_pool: ReqToTokenPool,
batch: ScheduleBatch,
@@ -47,7 +47,7 @@ from sglang.srt.speculative.adaptive_runtime_state import (
AdaptiveController,
SpecRuntimeState,
)
from sglang.srt.speculative.base_spec_worker import BaseDraftWorker, BaseSpecWorker
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase
from sglang.srt.speculative.draft_utils import DraftBackendFactory
from sglang.srt.speculative.eagle_draft_cuda_graph_runner import (
EAGLEDraftCudaGraphRunner,
@@ -117,7 +117,7 @@ def _get_plan_stream(
return None, contextlib.nullcontext()
class EagleDraftWorker(BaseDraftWorker):
class EagleDraftWorker(EagleDraftWorkerBase):
def __init__(
self,
server_args: ServerArgs,
@@ -408,7 +408,7 @@ class EagleDraftWorker(BaseDraftWorker):
def draft(self, batch: ScheduleBatch):
draft_input: EagleDraftInput = batch.spec_info
forward_batch, can_cuda_graph = draft_input.prepare_for_v2_draft(
forward_batch, can_cuda_graph = draft_input.prepare_for_draft(
self.req_to_token_pool,
batch,
self.cuda_graph_runner,
@@ -730,7 +730,8 @@ class EagleDraftWorker(BaseDraftWorker):
# Prepare for draft extend in a separate stream
with self.plan_stream_ctx:
forward_batch = draft_extend_input.prepare_for_extend_to_fill_draft_kvcache(
forward_batch = self.prepare_for_draft_extend(
draft_extend_input,
batch,
batch_result.next_token_ids,
self.speculative_num_draft_tokens,
@@ -1213,12 +1214,10 @@ class EAGLEWorkerV2(BaseSpecWorker):
# Batch 1: Target verify
# Prepare for target verify in a separate stream
with self.plan_stream_ctx:
verify_forward_batch, can_run_cuda_graph = (
verify_input.prepare_for_v2_verify(
self.req_to_token_pool,
batch,
self.target_worker,
)
verify_forward_batch, can_run_cuda_graph = verify_input.prepare_for_verify(
self.req_to_token_pool,
batch,
self.target_worker,
)
# Cover post-prepare rebinds: draft_token, plan_stream-allocated out_cache_loc.
@@ -1264,7 +1263,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
# Run target verify batch in the main compute stream (GPU compute).
# Metadata init is skipped iff cuda-graph already ran replay_prepare —
# prepare_for_v2_verify marked the batch in exactly that case; the
# 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(
@@ -1343,7 +1342,7 @@ class EAGLEWorkerV2(BaseSpecWorker):
# 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_extend_to_fill_draft_kvcache.
# 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,
@@ -41,7 +41,7 @@ from sglang.srt.model_executor.forward_batch_info import (
from sglang.srt.model_executor.forward_context import ForwardContext, forward_context
from sglang.srt.model_executor.pool_configurator import MemoryPoolConfig
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.base_spec_worker import BaseDraftWorker
from sglang.srt.speculative.base_spec_worker import EagleDraftWorkerBase
from sglang.srt.speculative.eagle_utils import (
build_tree_kernel_efficient,
organize_draft_results,
@@ -77,7 +77,7 @@ from sglang.srt.utils.async_probe import (
logger = logging.getLogger(__name__)
class FrozenKVMTPDraftWorker(BaseDraftWorker, TpModelWorker):
class FrozenKVMTPDraftWorker(EagleDraftWorkerBase, TpModelWorker):
"""Frozen-KV MTP draft worker.
The assistant reads target KV only. It reuses EAGLE's verify input/output
@@ -130,7 +130,7 @@ class FrozenKVMTPDraftWorker(BaseDraftWorker, TpModelWorker):
with (
empty_context()
), speculative_moe_backend_context(), speculative_moe_a2a_backend_context():
# NOTE: call TpModelWorker.__init__ explicitly -- BaseDraftWorker is
# NOTE: call TpModelWorker.__init__ explicitly -- EagleDraftWorkerBase is
# an ABC with no __init__, so cooperative super() would be ambiguous.
TpModelWorker.__init__(
self,
@@ -180,7 +180,7 @@ class FrozenKVMTPDraftWorker(BaseDraftWorker, TpModelWorker):
req_to_token_pool=None,
token_to_kv_pool_allocator=None,
):
# NOTE: call TpModelWorker explicitly -- BaseDraftWorker precedes it in
# NOTE: call TpModelWorker explicitly -- EagleDraftWorkerBase precedes it in
# the MRO and its alloc_memory_pool is a no-op stub.
TpModelWorker.alloc_memory_pool(
self,
@@ -583,7 +583,7 @@ class FrozenKVMTPDraftWorker(BaseDraftWorker, TpModelWorker):
)
def draft_extend(self):
# BaseDraftWorker contract. Frozen has no draft-KV extend forward; the
# EagleDraftWorkerBase contract. Frozen has no draft-KV extend forward; the
# orchestrator calls `_draft_extend_for_{prefill,decode}` directly.
pass
@@ -41,7 +41,7 @@ from sglang.srt.model_executor.forward_batch_info import (
ForwardBatch,
)
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.base_spec_worker import BaseDraftWorker, BaseSpecWorker
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase
from sglang.srt.speculative.draft_utils import DraftBackendFactory
from sglang.srt.speculative.eagle_info import (
EagleDraftExtendInput,
@@ -92,7 +92,7 @@ def _get_plan_stream(
return None, contextlib.nullcontext()
class MultiLayerEagleDraftWorker(BaseDraftWorker):
class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
def __init__(
self,
server_args: ServerArgs,
@@ -253,7 +253,7 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
def draft(self, batch: ScheduleBatch):
draft_input: EagleDraftInput = batch.spec_info
forward_batch, can_cuda_graph = draft_input.prepare_for_v2_draft(
forward_batch, can_cuda_graph = draft_input.prepare_for_draft(
self.req_to_token_pool,
batch,
self.cuda_graph_runner,
@@ -518,7 +518,8 @@ class MultiLayerEagleDraftWorker(BaseDraftWorker):
# Prepare for draft extend in a separate stream
# Notice that here we use batch_result.next_token_ids as the input ids
with self.plan_stream_ctx:
forward_batch = draft_extend_input.prepare_for_extend_to_fill_draft_kvcache(
forward_batch = self.prepare_for_draft_extend(
draft_extend_input,
batch,
batch_result.next_token_ids,
self.speculative_num_draft_tokens,
@@ -805,12 +806,10 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
# Batch 1: Target verify
# Prepare for target verify in a separate stream
with self.plan_stream_ctx:
verify_forward_batch, can_run_cuda_graph = (
verify_input.prepare_for_v2_verify(
self.req_to_token_pool,
batch,
self.target_worker,
)
verify_forward_batch, can_run_cuda_graph = verify_input.prepare_for_verify(
self.req_to_token_pool,
batch,
self.target_worker,
)
# Cover post-prepare rebinds: draft_token, plan_stream-allocated out_cache_loc.
@@ -834,7 +833,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
),
)
# NOTE: metadata init is skipped here unconditionally, although
# prepare_for_v2_verify only plans when cuda-graph replay_prepare ran.
# prepare_for_verify only plans when cuda-graph replay_prepare 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
@@ -12,7 +12,7 @@ from sglang.srt.managers.tp_worker import TpModelWorker
from sglang.srt.model_executor.forward_batch_info import ForwardMode
from sglang.srt.observability.req_time_stats import set_time_batch
from sglang.srt.server_args import ServerArgs
from sglang.srt.speculative.base_spec_worker import BaseDraftWorker, BaseSpecWorker
from sglang.srt.speculative.base_spec_worker import BaseSpecWorker, EagleDraftWorkerBase
from sglang.srt.speculative.cpp_ngram.ngram_corpus import NgramCorpus
from sglang.srt.speculative.ngram_info import NgramVerifyInput
from sglang.srt.speculative.spec_utils import (
@@ -110,7 +110,7 @@ class NGRAMWorker(BaseSpecWorker):
return self._target_worker
@property
def draft_worker(self) -> Optional[BaseDraftWorker]:
def draft_worker(self) -> Optional[EagleDraftWorkerBase]:
# NGRAM has no draft model; drafts come from the CPU-side corpus.
return None
+1 -1
View File
@@ -120,7 +120,7 @@ def record_stream_each(tensors, stream):
def record_stream_for_v2_verify(batch, verify_input, fwd_stream):
"""Mark pre-prepare SB / verify_input GPU tensors as used on `fwd_stream`.
Spec V2 mutates SB mid-forward (`prepare_for_v2_verify` rebinds
Spec V2 mutates SB mid-forward (`prepare_for_verify` rebinds
`batch.input_ids` / `out_cache_loc`; `_draft_extend_for_decode` later
replaces `batch.input_ids` again). Each rebind drops the only SB Python
ref to the old tensor while the verify forward kernel may still be
@@ -154,7 +154,7 @@ def _scan_srt():
def _draft_worker_classes():
"""All transitive BaseDraftWorker subclasses under speculative/."""
"""All transitive EagleDraftWorkerBase subclasses under speculative/."""
by_name = {}
for path in sorted(_SPECULATIVE_DIR.glob("*.py")):
rel = path.relative_to(_SRT_DIR).as_posix()
@@ -166,7 +166,7 @@ def _draft_worker_classes():
}
by_name[node.name] = (rel, node, bases)
workers = {"BaseDraftWorker"}
workers = {"EagleDraftWorkerBase"}
changed = True
while changed:
changed = False
@@ -177,7 +177,7 @@ def _draft_worker_classes():
return [
(rel, node)
for name, (rel, node, _) in sorted(by_name.items())
if name in workers and name != "BaseDraftWorker"
if name in workers and name != "EagleDraftWorkerBase"
]