[Spec] Stage Inkling MTP draft metadata before verify (#38169)

Co-authored-by: Qiaolin-Yu <liin1211@outlook.com>
This commit is contained in:
paulzhang-tm
2026-09-10 15:32:19 -07:00
committed by GitHub
co-authored by Qiaolin-Yu
parent a63efd9056
commit bb15be6d79
5 changed files with 205 additions and 35 deletions
@@ -129,6 +129,8 @@ class AttentionBackend(ABC):
Default: no-op.
"""
supports_draft_extend_metadata_staging: bool = False
def draft_extend_metadata_captured_in_graph(self) -> bool:
"""True when :py:meth:`init_forward_metadata_in_graph` fully rebuilds
this backend's DRAFT_EXTEND_V2 replay metadata inside the captured
@@ -451,6 +451,20 @@ class FlashAttentionBackend(AttentionBackend):
),
)
@property
def supports_draft_extend_metadata_staging(self) -> bool:
return (
self.topk == 1
and not self.kv_index_translator.is_translating
and self.draft_extend_metadata_captured_in_graph()
)
def stage_draft_extend_metadata(self, forward_batch: ForwardBatch):
self.forward_metadata = self.draft_extend_metadata[forward_batch.batch_size]
self.forward_metadata.max_seq_len_k = self.max_context_len
self.forward_metadata_spec_decode_expand = None
self.init_forward_metadata_in_graph(forward_batch)
def _in_graph_full_to_swa_index_mapping(self) -> Optional[torch.Tensor]:
# The in-graph SWA translation needs the raw mapping tensor; v2p-table
# pools (UnifiedSWAKVPool) keep it None and must stay on the eager
@@ -501,16 +501,22 @@ class InklingShortConvAttnBackend(ShortConvAttnBackend):
mamba_track_indices: Optional[torch.Tensor],
mamba_steps_to_track: Optional[torch.Tensor],
) -> None:
"""Commit the TARGET_VERIFY conv windows at each request's last accepted step.
Slot ids come from ``req_pool_indices``, not the per-step
``self._cache_indices``: this runs after the forward context exits, so that
buffer may already belong to a later forward.
"""
"""Commit the TARGET_VERIFY conv windows at each request's last accepted step."""
pool = self.req_to_token_pool
bs = req_pool_indices.shape[0]
if self._slot_gather_recordable:
assert (
self._cache_indices_buf is not None
and self._cache_indices_buf.shape[0] >= bs
)
slot_ids = self._cache_indices_buf[:bs]
else:
slot_ids = self._translate_mamba_indices(
pool.get_mamba_indices(req_pool_indices)
)
scatter_mamba_states_after_mtp_verify(
pool.get_speculative_mamba2_params_all_layers(),
self._translate_mamba_indices(pool.get_mamba_indices(req_pool_indices)),
slot_ids,
last_correct_step_indices,
mamba_track_indices,
mamba_steps_to_track,
@@ -548,6 +554,43 @@ class InklingShortConvHybridAttnBackend(ShortConvHybridAttnBackend):
is Inkling's own, not the generic mamba scatter.
"""
@property
def supports_draft_extend_metadata_staging(self) -> bool:
return (
self.full_attn_backend.supports_draft_extend_metadata_staging
and self.short_conv_backend._slot_gather_recordable
)
def init_forward_metadata_out_graph(
self, forward_batch: ForwardBatch, in_capture: bool = False
):
if (
forward_batch.forward_mode.is_draft_extend_v2()
and self.supports_draft_extend_metadata_staging
):
if in_capture:
self.full_attn_backend.init_forward_metadata_out_graph(
forward_batch, in_capture=True
)
self.full_attn_backend.stage_draft_extend_metadata(forward_batch)
self.short_conv_backend._prepare_slot_indices(forward_batch)
else:
super().init_forward_metadata_out_graph(
forward_batch, in_capture=in_capture
)
def init_forward_metadata_in_graph(self, forward_batch: ForwardBatch):
if (
forward_batch.forward_mode.is_draft_extend_v2()
and self.supports_draft_extend_metadata_staging
):
self.short_conv_backend._reset_step_state()
self.short_conv_backend._refresh_sconv_metadata(
forward_batch, on_graph_path=True
)
else:
super().init_forward_metadata_in_graph(forward_batch)
def sconv_state(self, *, layer_id: int, stream: int) -> torch.Tensor:
return self.short_conv_backend.sconv_state(layer_id=layer_id, stream=stream)
@@ -610,4 +653,7 @@ class InklingShortConvHybridAttnBackend(ShortConvHybridAttnBackend):
)
def draft_extend_metadata_captured_in_graph(self) -> bool:
return self.full_attn_backend.draft_extend_metadata_captured_in_graph()
return (
not self.supports_draft_extend_metadata_staging
and self.full_attn_backend.draft_extend_metadata_captured_in_graph()
)
@@ -493,7 +493,10 @@ class MultiLayerEagleDraftExtendCudaGraphRunner(DecodeCudaGraphRunner):
out_cache_loc=buffers.out_cache_loc[:num_tokens],
spec_info=spec_info,
)
if not self.metadata_captured_in_graph:
if (
not self.metadata_captured_in_graph
and not self.attn_backend.supports_draft_extend_metadata_staging
):
self.eagle_worker.draft_extend_attn_backend_list[
self.step
].init_forward_metadata_out_graph(fb_view)
@@ -712,7 +715,47 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
def _prepare_extra(self, forward_batch: ForwardBatch) -> None:
"""Hook for subclasses to populate extra per-call buffers (e.g. sconv)."""
def prepare(self, forward_batch: ForwardBatch):
def stage_shared_reads(
self, *, seq_lens, req_pool_indices, out_cache_loc, positions=None
):
raw_bs = req_pool_indices.shape[0]
bs = self.get_runner(0)._pad_to_bucket(raw_bs, self.capture_bs)
buffers = self.buffers
buffers.seq_lens[:bs].fill_(self.seq_len_fill_value)
buffers.seq_lens[:raw_bs].copy_(seq_lens)
buffers.req_pool_indices[:bs].zero_()
buffers.req_pool_indices[:raw_bs].copy_(req_pool_indices)
num_tokens = raw_bs * self.captured_req_width
buffers.out_cache_loc[: bs * self.captured_req_width].zero_()
buffers.out_cache_loc[:num_tokens].copy_(out_cache_loc)
if positions is not None:
buffers.positions[:num_tokens].copy_(positions)
self._stage_metadata(bs, raw_bs)
self._staged_bs = bs
def _stage_metadata(self, bs: int, raw_bs: int):
backends = [
b
for b in self.draft_extend_attn_backend_list
if b.supports_draft_extend_metadata_staging
and not b.draft_extend_metadata_captured_in_graph()
]
if not backends:
return
buffers = self.buffers
buffers.req_pool_indices[raw_bs:bs].zero_()
batch = SimpleNamespace(
batch_size=bs,
forward_mode=ForwardMode.DRAFT_EXTEND_V2,
req_pool_indices=buffers.req_pool_indices[:bs],
seq_lens=buffers.seq_lens[:bs],
extend_seq_lens=buffers.extend_seq_lens[:bs],
out_cache_loc=buffers.out_cache_loc[: bs * self.captured_req_width],
)
for backend in backends:
backend.init_forward_metadata_out_graph(batch)
def prepare(self, forward_batch: ForwardBatch, *, staged: bool = False):
"""Populate the shared buffers once from ``forward_batch`` and bucketize
the batch size. Subsequent ``replay(step)`` calls reuse this state."""
buffers = self.buffers
@@ -801,6 +844,10 @@ class MultiLayerEagleMultiStepDraftExtendCudaGraphRunner:
seq_lens_sum = seq_lens_sum + (bs - raw_bs) * self.seq_len_fill_value
self.seq_lens_sum = seq_lens_sum
if staged:
assert bs == self._staged_bs
else:
self._stage_metadata(bs, raw_bs)
self._prepare_extra(forward_batch)
def replay(self, step: int):
@@ -852,10 +899,9 @@ class OneGraphMultiLayerEagleMultiStepDraftExtendCudaGraphRunner(
forwards + the inter-step input_ids rotation in ONE graph per bucket, instead
of one graph per step. The worker drops its per-step rotation (rotates_in_graph).
Each step's replay metadata is emitted in-graph via
init_forward_metadata_in_graph (no Python may run between
captured steps). seq_lens / req_pool_indices / extend_seq_lens are chain-constant
(only input_ids rotates), so per-step in-graph metadata is correct.
Each step refreshes metadata in-graph or stages it before replay; no Python
may run between captured steps. seq_lens / req_pool_indices / extend_seq_lens
are chain-constant (only input_ids rotates).
Rejection sampling is supported by sampling X ~ q inside the graph
(_sample_draft_proposal, selected by the draft_probs buffer's presence):
@@ -114,6 +114,8 @@ logger = logging.getLogger(__name__)
class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
last_draft_extend_staged: bool = False
def __init__(
self,
server_args: ServerArgs,
@@ -289,8 +291,8 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
def _compute_boundary_kv_locs_positions(self, batch):
if self.draft_extend_num_front_tokens == 0 or batch.forward_mode.is_idle():
return None, None, None
locs, positions = compute_widened_draft_extend_locs_positions(
return None, None
return compute_widened_draft_extend_locs_positions(
batch.seq_lens,
batch.req_pool_indices,
self.req_to_token_pool.req_to_token,
@@ -299,11 +301,6 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
self.draft_extend_num_front_tokens,
self.draft_extend_num_warmup_tokens,
)
ready_event = None
if self.plan_stream:
ready_event = torch.get_device_module(self.device).Event()
ready_event.record()
return locs, positions, ready_event
def _seed_boundary_kv_stash(self, forward_batch, target_hidden_states):
if (
@@ -411,15 +408,11 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
before_mem = get_available_gpu_memory(self.device, self.gpu_id)
if not _is_npu:
# The single-CG runner replays with no Python between steps, so the
# attn backend must fully rebuild its per-step metadata as captured
# tensor ops; anything less gets capture-time-stale metadata (e.g.
# SWA translations, which only the eager replay path refreshes).
# Per-depth pools (banded MTP) mean per-depth backends — EVERY step
# must satisfy this, not just step 0.
# Every step must refresh metadata in-graph or stage it before replay.
draft_backend = self.draft_runner_list[0].attn_backend
backend_supports_single_cg = all(
runner.attn_backend.draft_extend_metadata_captured_in_graph()
or runner.attn_backend.supports_draft_extend_metadata_staging
for runner in self.draft_runner_list
)
if envs.SGLANG_ENABLE_SINGLE_CG_DRAFT.get() and backend_supports_single_cg:
@@ -429,8 +422,8 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
else:
if envs.SGLANG_ENABLE_SINGLE_CG_DRAFT.get():
logger.warning(
"SGLANG_ENABLE_SINGLE_CG_DRAFT is on but %s does not fully "
"rebuild its draft-extend metadata in-graph; falling back "
"SGLANG_ENABLE_SINGLE_CG_DRAFT is on but %s cannot refresh "
"its draft-extend metadata for a combined graph; falling back "
"to per-step draft graphs.",
type(draft_backend).__name__,
)
@@ -724,8 +717,57 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
return next_draft_input
def _draft_extend_plan_for_decode(self, batch: ScheduleBatch) -> bool:
runner = self.cuda_graph_runner_for_draft_extend
if runner is None:
return False
assert self.draft_extend_attn_backend_list, (
"Draft graphs require initialized attention backends"
)
target_graph = self.target_worker.model_runner.decode_cuda_graph_runner
# In-graph mapping reads still require the final draft fence.
if (
batch.forward_mode.is_idle()
or getattr(target_graph, "in_graph_metadata_prep_done", None) is None
or not all(
b.supports_draft_extend_metadata_staging
and not b.draft_extend_metadata_captured_in_graph()
for b in self.draft_extend_attn_backend_list
)
or runner.require_mlp_tp_gather
or len(batch.seq_lens) > runner.max_bs
or (runner.disable_padding and len(batch.seq_lens) not in runner.capture_bs)
):
return False
assert self.topk == 1, "Draft-extend metadata staging requires topk=1"
locs, positions = self._compute_boundary_kv_locs_positions(batch)
if locs is None:
from sglang.kernels.ops.speculative.cache_locs import (
assign_extend_cache_locs_uniform_func,
)
locs = assign_extend_cache_locs_uniform_func(
req_pool_indices=batch.req_pool_indices,
req_to_token=self.req_to_token_pool.req_to_token,
start_offset=batch.seq_lens,
batch_size=len(batch.seq_lens),
draft_token_num=self.speculative_num_draft_tokens,
device=batch.device,
)
runner.stage_shared_reads(
seq_lens=batch.seq_lens + self.speculative_num_draft_tokens,
req_pool_indices=batch.req_pool_indices,
out_cache_loc=locs,
positions=positions,
)
return True
def _draft_extend_for_decode(
self, batch: ScheduleBatch, batch_result: GenerationBatchResult
self,
batch: ScheduleBatch,
batch_result: GenerationBatchResult,
*,
staged: bool = False,
):
# Batch 2: Draft extend
draft_extend_input = EagleDraftExtendInput(
@@ -742,9 +784,19 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
# Prepare for draft extend in a separate stream
# Notice that here we use batch_result.next_token_ids as the input ids
boundary_kv_locs, boundary_kv_positions, boundary_kv_ready_event = (
self._compute_boundary_kv_locs_positions(batch)
)
if staged and self.draft_extend_num_front_tokens:
runner = self.cuda_graph_runner_for_draft_extend
num_tokens = len(batch.seq_lens) * runner.captured_req_width
boundary_kv_locs = runner.buffers.out_cache_loc[:num_tokens]
boundary_kv_positions = runner.buffers.positions[:num_tokens]
else:
boundary_kv_locs, boundary_kv_positions = (
self._compute_boundary_kv_locs_positions(batch)
)
boundary_kv_ready_event = None
if boundary_kv_locs is not None and self.plan_stream:
boundary_kv_ready_event = torch.get_device_module(self.device).Event()
boundary_kv_ready_event.record()
with self.plan_stream_ctx:
if boundary_kv_ready_event is not None:
@@ -798,7 +850,7 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
cgr = self.cuda_graph_runner_for_draft_extend
# Populate the single shared buffer set once; each step replays
# against it and the chain is advanced in place between steps.
cgr.prepare(forward_batch)
cgr.prepare(forward_batch, staged=staged)
rotates_in_graph = cgr.rotates_in_graph
for step in range(self.speculative_num_steps):
_out, ret_topk_p, ret_topk_index = cgr.replay(step)
@@ -936,6 +988,10 @@ class MultiLayerEagleDraftWorker(EagleDraftWorkerBase):
ret_draft_probs = torch.stack(ret_draft_probs_list, dim=1)
next_draft_input.draft_probs = ret_draft_probs
self.last_draft_extend_staged = bool(
staged and can_run_decode_cuda_graph and batch_result.can_run_cuda_graph
)
class MultiLayerEagleWorkerV2(BaseSpecWorker):
def __init__(
@@ -979,6 +1035,8 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
@property
def last_shared_read_runner(self):
if self.draft_worker.last_draft_extend_staged:
return self.target_worker.model_runner
# Multi-layer eagle has no draft forward, only draft extend.
return self._draft_worker.draft_runner
@@ -998,6 +1056,7 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
def forward_batch_generation(
self, batch: ScheduleBatch, on_publish=None, grammar_barrier=None
):
self.draft_worker.last_draft_extend_staged = False
if batch.forward_mode.is_extend() or batch.is_extend_in_batch:
# Target prefill
target_capture_mode = (
@@ -1044,11 +1103,14 @@ class MultiLayerEagleWorkerV2(BaseSpecWorker):
verify_input: EagleVerifyInput = self.draft_worker.draft(batch)
assert verify_input.is_verify_input()
batch.spec_info = verify_input
staged = self.draft_worker._draft_extend_plan_for_decode(batch)
batch_output = self.verify(batch, grammar_barrier=grammar_barrier)
# Publish before draft_extend so the fence is at verify-end.
if on_publish is not None:
on_publish(batch_output.new_seq_lens)
self.draft_worker._draft_extend_for_decode(batch, batch_output)
self.draft_worker._draft_extend_for_decode(
batch, batch_output, staged=staged
)
return batch_output
def verify(self, batch: ScheduleBatch, grammar_barrier=None):