From bb15be6d79d3b127295b0436b5164757335c2d1a Mon Sep 17 00:00:00 2001 From: paulzhang-tm Date: Thu, 10 Sep 2026 18:32:19 -0400 Subject: [PATCH] [Spec] Stage Inkling MTP draft metadata before verify (#38169) Co-authored-by: Qiaolin-Yu --- .../srt/layers/attention/base_attn_backend.py | 2 + .../attention/flashattention_backend.py | 14 +++ .../attention/linear/inkling_sconv_backend.py | 62 +++++++++-- ...er_eagle_draft_extend_cuda_graph_runner.py | 58 +++++++++- .../multi_layer_eagle_worker_v2.py | 104 ++++++++++++++---- 5 files changed, 205 insertions(+), 35 deletions(-) diff --git a/python/sglang/srt/layers/attention/base_attn_backend.py b/python/sglang/srt/layers/attention/base_attn_backend.py index b988c9a55..5f750ee9e 100644 --- a/python/sglang/srt/layers/attention/base_attn_backend.py +++ b/python/sglang/srt/layers/attention/base_attn_backend.py @@ -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 diff --git a/python/sglang/srt/layers/attention/flashattention_backend.py b/python/sglang/srt/layers/attention/flashattention_backend.py index 79d100be9..01e78b7ee 100644 --- a/python/sglang/srt/layers/attention/flashattention_backend.py +++ b/python/sglang/srt/layers/attention/flashattention_backend.py @@ -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 diff --git a/python/sglang/srt/layers/attention/linear/inkling_sconv_backend.py b/python/sglang/srt/layers/attention/linear/inkling_sconv_backend.py index 8948a22ca..8af6d78f1 100644 --- a/python/sglang/srt/layers/attention/linear/inkling_sconv_backend.py +++ b/python/sglang/srt/layers/attention/linear/inkling_sconv_backend.py @@ -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() + ) diff --git a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py index 77f2f0b66..dc0d0230f 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_draft_extend_cuda_graph_runner.py @@ -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): diff --git a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py index f5d619894..b62755c53 100644 --- a/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py +++ b/python/sglang/srt/speculative/multi_layer_eagle_worker_v2.py @@ -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):