[Spec] Stage Inkling MTP draft metadata before verify (#38169)
Co-authored-by: Qiaolin-Yu <liin1211@outlook.com>
This commit is contained in:
co-authored by
Qiaolin-Yu
parent
a63efd9056
commit
bb15be6d79
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user