From 48b3b82f1cabb3a62937747653d82756fae16056 Mon Sep 17 00:00:00 2001 From: Yizhou Liu Date: Sun, 26 Jul 2026 07:54:45 +0000 Subject: [PATCH 01/11] feat(spec): refresh DeepSeek V4 SWA metadata inside the draft loop using backend builder and gate fused draft decode graphs by backend capability Signed-off-by: Yizhou Liu --- .../test_gpu_autoregressive_speculator.py | 51 ++++++++ vllm/config/speculative.py | 4 + vllm/v1/attention/backend.py | 7 + vllm/v1/attention/backends/mla/sparse_swa.py | 34 +++++ .../spec_decode/autoregressive/speculator.py | 120 +++++++++++++++++- vllm/v1/worker/utils.py | 11 ++ 6 files changed, 221 insertions(+), 6 deletions(-) diff --git a/tests/v1/worker/test_gpu_autoregressive_speculator.py b/tests/v1/worker/test_gpu_autoregressive_speculator.py index 9f5f90a9775e..7f4556c4f858 100644 --- a/tests/v1/worker/test_gpu_autoregressive_speculator.py +++ b/tests/v1/worker/test_gpu_autoregressive_speculator.py @@ -3,6 +3,7 @@ from contextlib import nullcontext from types import SimpleNamespace +from unittest.mock import Mock import pytest import torch @@ -16,6 +17,7 @@ from vllm.model_executor.models.mistral_large_3_eagle import ( EagleMistralLarge3ForCausalLM, ) +from vllm.v1.worker.gpu.cudagraph_utils import BatchExecutionDescriptor from vllm.v1.worker.gpu.spec_decode import speculator as base_spec_module from vllm.v1.worker.gpu.spec_decode.autoregressive import speculator as spec_module from vllm.v1.worker.gpu.spec_decode.autoregressive.speculator import ( @@ -278,3 +280,52 @@ def test_run_model_reuses_tensor_return_for_mtp(monkeypatch): assert actual_logits_hidden is hidden assert actual_feedback_hidden is hidden + + +@pytest.mark.parametrize( + ( + "cg_mode", + "use_fused_decode_graph", + "expected_eager_calls", + "expected_graph_replays", + ), + [ + (CUDAGraphMode.NONE, True, 3, 0), + (CUDAGraphMode.FULL, True, 0, 1), + (CUDAGraphMode.FULL, False, 0, 3), + ], +) +def test_multi_step_decode_replays_captured_graph_as_expected( + cg_mode, + use_fused_decode_graph, + expected_eager_calls, + expected_graph_replays, +): + speculator = object.__new__(_TestSpeculator) + speculator.num_speculative_steps = 4 + speculator.current_draft_step = torch.tensor(0) + speculator.input_buffers = SimpleNamespace( + positions=torch.arange(2), + query_start_loc=torch.arange(3), + ) + speculator.idx_mapping = torch.arange(2) + speculator.use_fused_decode_graph = use_fused_decode_graph + generate_draft = Mock() + speculator._generate_draft = generate_draft + run_fullgraph = Mock() + speculator.decode_cudagraph_manager = SimpleNamespace(run_fullgraph=run_fullgraph) + batch_desc = BatchExecutionDescriptor( + cg_mode=cg_mode, + num_tokens=2, + num_reqs=2, + ) + + speculator._multi_step_decode( + num_reqs=2, + skip_attn=True, + batch_desc=batch_desc, + num_tokens_across_dp=None, + ) + + assert generate_draft.call_count == expected_eager_calls + assert run_fullgraph.call_count == expected_graph_replays diff --git a/vllm/config/speculative.py b/vllm/config/speculative.py index fa8c282007b2..1a817e759adc 100644 --- a/vllm/config/speculative.py +++ b/vllm/config/speculative.py @@ -138,6 +138,9 @@ class SpeculativeConfig: will use the default version.""" # Advanced control + enable_fused_decode_graph: bool = True + """Fuse autoregressive draft decode steps into one CUDA graph when the + attention backends support it. Unsupported backends use per-step graphs.""" disable_padded_drafter_batch: bool = False """Disable input padding for speculative decoding. If set to True, speculative input batches can contain sequences of different lengths, @@ -330,6 +333,7 @@ def compute_hash(self) -> str: # Convert to tuple to make it hashable factors.append(tuple(layer_ids)) + factors.append(self.enable_fused_decode_graph) hash_str = safe_hash(str(factors).encode(), usedforsecurity=False).hexdigest() return hash_str diff --git a/vllm/v1/attention/backend.py b/vllm/v1/attention/backend.py index d4698520f17c..25a7135c8aef 100644 --- a/vllm/v1/attention/backend.py +++ b/vllm/v1/attention/backend.py @@ -633,6 +633,10 @@ class AttentionMetadataBuilder(ABC, Generic[M]): supports_update_block_table: bool = False # Whether the builder constructor requires the block-table width. requires_block_table_width: ClassVar[bool] = False + # Does this backend support capture mutiple draft decode steps into one + # CUDA Graph (default: no), which requires no step-dependent metadata + # or can refresh metadata over steps. + supports_fused_decode_graph: bool = False @abstractmethod def __init__( @@ -757,6 +761,9 @@ def build_for_drafting( fast_build=True, ) + def refresh_meta_for_draft_decodes(self, metadata: M) -> None: + pass + def use_cascade_attention( self, common_prefix_len: int, diff --git a/vllm/v1/attention/backends/mla/sparse_swa.py b/vllm/v1/attention/backends/mla/sparse_swa.py index c823f82be29b..d7abcc07fcf0 100644 --- a/vllm/v1/attention/backends/mla/sparse_swa.py +++ b/vllm/v1/attention/backends/mla/sparse_swa.py @@ -400,6 +400,7 @@ class DeepseekSparseSWAMetadataBuilder(AttentionMetadataBuilder): reorder_batch_threshold: int | None = None _cudagraph_support: ClassVar[AttentionCGSupport] = AttentionCGSupport.UNIFORM_BATCH + supports_fused_decode_graph = True def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) @@ -649,6 +650,39 @@ def build( **deepseek_v4_fields, # type: ignore[arg-type] ) + def refresh_meta_for_draft_decodes( + self, + metadata: DeepseekSparseSWAMetadata, + ) -> None: + if metadata.num_decode_tokens == 0: + return + assert metadata.query_start_loc is not None + assert metadata.seq_lens is not None + assert metadata.token_to_req_indices is not None + assert metadata.is_valid_token is not None + assert metadata.decode_swa_indices is not None + assert metadata.decode_swa_lens is not None + + _compute_swa_indices_and_lens_kernel[(metadata.num_decode_tokens,)]( + metadata.decode_swa_indices, + metadata.decode_swa_indices.stride(0), + metadata.decode_swa_lens, + metadata.decode_swa_indices.shape[-1], + metadata.query_start_loc, + metadata.seq_lens, + metadata.token_to_req_indices, + metadata.is_valid_token, + metadata.block_table, + metadata.block_table.stride(0), + self.block_size, + token_offset=0, + TRITON_BLOCK_SIZE=1024, + ) + tile_sched = self.build_tile_scheduler(metadata.num_decode_tokens) + metadata.tile_sched_swaonly = tile_sched[_LAYER_TYPE_SWAONLY] + metadata.tile_sched_c4a = tile_sched[_LAYER_TYPE_C4A] + metadata.tile_sched_c128a = tile_sched[_LAYER_TYPE_C128A] + def build_tile_scheduler( self, num_decode_tokens: int ) -> dict[str, FlashMLASchedMeta | None]: diff --git a/vllm/v1/worker/gpu/spec_decode/autoregressive/speculator.py b/vllm/v1/worker/gpu/spec_decode/autoregressive/speculator.py index 5ceb6c7558d6..f96310d49c38 100644 --- a/vllm/v1/worker/gpu/spec_decode/autoregressive/speculator.py +++ b/vllm/v1/worker/gpu/spec_decode/autoregressive/speculator.py @@ -10,17 +10,21 @@ from vllm.forward_context import BatchDescriptor, set_forward_context from vllm.logger import init_logger from vllm.triton_utils import tl, triton +from vllm.v1.kv_cache_interface import KVCacheConfig from vllm.v1.worker.gpu.attn_utils import build_slot_mappings_by_layer +from vllm.v1.worker.gpu.block_table import BlockTables from vllm.v1.worker.gpu.cudagraph_utils import ( BatchExecutionDescriptor, get_uniform_token_count, ) from vllm.v1.worker.gpu.dp_utils import dispatch_cg_and_sync_dp from vllm.v1.worker.gpu.input_batch import InputBatch, InputBuffers +from vllm.v1.worker.gpu.model_states.interface import ModelState from vllm.v1.worker.gpu.spec_decode.autoregressive.cudagraph_utils import ( SpeculatorCudaGraphManager, ) from vllm.v1.worker.gpu.spec_decode.speculator import DraftModelSpeculator +from vllm.v1.worker.utils import AttentionGroup logger = init_logger(__name__) @@ -41,6 +45,7 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device): self.prefill_cudagraph_manager: SpeculatorCudaGraphManager | None = None self.decode_cudagraph_manager: SpeculatorCudaGraphManager | None = None + self.use_fused_decode_graph = False def load_model(self, target_model: nn.Module) -> None: super().load_model(target_model) @@ -76,6 +81,48 @@ def advance_draft_positions(self) -> bool: """ return True + def set_attn( + self, + model_state: ModelState, + kv_cache_config: KVCacheConfig, + block_tables: BlockTables, + target_input_buffers: InputBuffers, + target_attn_groups: list[list[AttentionGroup]], + ) -> None: + super().set_attn( + model_state, + kv_cache_config, + block_tables, + target_input_buffers, + target_attn_groups, + ) + self._configure_fused_decode_graph() + + def _configure_fused_decode_graph(self) -> None: + if ( + not self.speculative_config.enable_fused_decode_graph + or self.num_speculative_steps == 1 + ): + self.use_fused_decode_graph = False + return + + unsupported_backends = sorted( + { + attn_group.backend.get_name() + for attn_groups in self.attn_groups + for attn_group in attn_groups + if not attn_group.supports_fused_decode_graph + } + ) + self.use_fused_decode_graph = not unsupported_backends + if unsupported_backends: + logger.info_once( + "Fused draft decode graph is not supported by attention " + "backend(s) %s; falling back to rebuilding attention metadata " + "between draft steps.", + ", ".join(unsupported_backends), + ) + def init_cudagraph_manager(self, cudagraph_mode: CUDAGraphMode) -> None: # Initialize cudagraph manager for draft prefill (draft position 0). self.prefill_cudagraph_manager = SpeculatorCudaGraphManager( @@ -133,12 +180,15 @@ def capture(self) -> None: return self.on_multi_step_decode_begin(self.max_num_reqs) - # Capture the decode draft generation routine (model forward + - # sample + update_draft_inputs) for a single - # step. + # Capture either the fused decode loop or one decode step per graph. assert self.decode_cudagraph_manager is not None + decode_fn = ( + self._run_draft_decode_loop + if self.use_fused_decode_graph + else self._generate_draft + ) self.decode_cudagraph_manager.capture( - self._generate_draft, + decode_fn, self.model_state, self.input_buffers, self.block_tables, @@ -412,6 +462,27 @@ def _multi_step_decode( query_start_loc = self.input_buffers.query_start_loc[: num_reqs + 1] idx_mapping = self.idx_mapping[:num_reqs] + if batch_desc.cg_mode == CUDAGraphMode.FULL and self.use_fused_decode_graph: + if not skip_attn: + self.block_tables.compute_slot_mappings( + idx_mapping, + query_start_loc, + positions, + batch_desc.num_tokens, + ) + # Continuous draft decode replays usually consume only device metadata, + # so the host-side vars should have no effect and therefore pinned. + self._build_draft_attn_metadata( + num_reqs=num_reqs, + num_reqs_padded=batch_desc.num_reqs or num_reqs, + num_tokens_padded=batch_desc.num_tokens, + seq_lens_cpu_upper_bound=seq_lens_cpu_upper_bound, + step=1, + ) + assert self.decode_cudagraph_manager is not None + self.decode_cudagraph_manager.run_fullgraph(batch_desc) + return + attn_metadata = None slot_mappings_by_layer = None for step in range(1, self.num_speculative_steps): @@ -435,10 +506,8 @@ def _multi_step_decode( step=step, ) - # Update the current draft step. self.current_draft_step.fill_(step) - # Generate draft tokens for the current step. if batch_desc.cg_mode == CUDAGraphMode.FULL: assert self.decode_cudagraph_manager is not None self.decode_cudagraph_manager.run_fullgraph(batch_desc) @@ -452,6 +521,45 @@ def _multi_step_decode( cudagraph_runtime_mode=batch_desc.cg_mode, ) + def _run_draft_decode_loop( + self, + num_reqs: int, + num_tokens_padded: int, + attn_metadata: dict[str, Any] | None, + slot_mappings: dict[str, torch.Tensor] | None, + num_tokens_across_dp: torch.Tensor | None, + cudagraph_runtime_mode: CUDAGraphMode = CUDAGraphMode.NONE, + ) -> None: + idx_mapping = self.idx_mapping[:num_reqs] + positions = self.input_buffers.positions[:num_reqs] + query_start_loc = self.input_buffers.query_start_loc[: num_reqs + 1] + attn_groups = ( + [group for groups in self.attn_groups for group in groups] + if attn_metadata is not None + else [] + ) + + for step in range(1, self.num_speculative_steps): + self.current_draft_step.fill_(step) + self._generate_draft( + num_reqs, + num_tokens_padded, + attn_metadata, + slot_mappings, + num_tokens_across_dp, + cudagraph_runtime_mode, + ) + if step < self.num_speculative_steps - 1 and attn_metadata is not None: + if self.advance_draft_positions: + self.block_tables.compute_slot_mappings( + idx_mapping, + query_start_loc, + positions, + num_tokens_padded, + ) + for attn_group in attn_groups: + attn_group.refresh_meta_for_draft_decodes(attn_metadata) + def _generate_draft( self, num_reqs: int, diff --git a/vllm/v1/worker/utils.py b/vllm/v1/worker/utils.py index 9d056cc67c24..f0b12cda238c 100644 --- a/vllm/v1/worker/utils.py +++ b/vllm/v1/worker/utils.py @@ -299,6 +299,17 @@ def get_metadata_builder(self, ubatch_id: int = 0) -> AttentionMetadataBuilder: assert len(self.metadata_builders) > ubatch_id return self.metadata_builders[ubatch_id] + @property + def supports_fused_decode_graph(self) -> bool: + return self.get_metadata_builder().supports_fused_decode_graph + + def refresh_meta_for_draft_decodes( + self, + attn_metadata: Mapping[str, Any], + ) -> None: + metadata = attn_metadata[self.layer_names[0]] + self.get_metadata_builder().refresh_meta_for_draft_decodes(metadata) + def select_common_block_size( kv_manager_block_size: int, From b9b1a9a51116818d6fc4ab655c2ca3e521607570 Mon Sep 17 00:00:00 2001 From: Yizhou Liu Date: Sun, 26 Jul 2026 09:22:43 +0000 Subject: [PATCH 02/11] feat(spec): enable fused draft decode graphs for FA3 without DCP Signed-off-by: Yizhou Liu --- .../test_gpu_autoregressive_speculator.py | 118 +++++++++++++ vllm/v1/attention/backends/flash_attn.py | 159 +++++++++++++----- 2 files changed, 235 insertions(+), 42 deletions(-) diff --git a/tests/v1/worker/test_gpu_autoregressive_speculator.py b/tests/v1/worker/test_gpu_autoregressive_speculator.py index 7f4556c4f858..bff851592395 100644 --- a/tests/v1/worker/test_gpu_autoregressive_speculator.py +++ b/tests/v1/worker/test_gpu_autoregressive_speculator.py @@ -17,6 +17,8 @@ from vllm.model_executor.models.mistral_large_3_eagle import ( EagleMistralLarge3ForCausalLM, ) +from vllm.v1.attention.backends import flash_attn as flash_attn_module +from vllm.v1.attention.backends.flash_attn import FlashAttentionMetadata from vllm.v1.worker.gpu.cudagraph_utils import BatchExecutionDescriptor from vllm.v1.worker.gpu.spec_decode import speculator as base_spec_module from vllm.v1.worker.gpu.spec_decode.autoregressive import speculator as spec_module @@ -324,8 +326,124 @@ def test_multi_step_decode_replays_captured_graph_as_expected( num_reqs=2, skip_attn=True, batch_desc=batch_desc, + seq_lens_cpu_upper_bound=None, num_tokens_across_dp=None, ) assert generate_draft.call_count == expected_eager_calls assert run_fullgraph.call_count == expected_graph_replays + + +def test_refresh_meta_for_draft_decodes_updates_fa3_scheduler_metadata( + monkeypatch, +): + builder = object.__new__(flash_attn_module.FlashAttentionMetadataBuilder) + builder.aot_schedule = True + builder.use_full_cuda_graph = True + builder.scheduler_metadata = torch.zeros(8, dtype=torch.int32) + builder.cache_config = SimpleNamespace(cache_dtype="bfloat16") + builder.kv_cache_dtype = torch.bfloat16 + builder.num_heads_q = 2 + builder.num_heads_kv = 1 + builder.headdim = 128 + builder.block_size = 16 + builder.dcp_world_size = 1 + builder.dcp_rank = 0 + builder.cp_kv_cache_interleave_size = 1 + builder.aot_sliding_window = None + + expected = torch.tensor([7, 8, 9], dtype=torch.int32) + + def fake_get_scheduler_metadata(**kwargs): + return expected + + monkeypatch.setattr(builder, "_get_scheduler_metadata", fake_get_scheduler_metadata) + + metadata = FlashAttentionMetadata( + num_actual_tokens=3, + max_query_len=2, + query_start_loc=torch.tensor([0, 1, 3], dtype=torch.int32), + max_seq_len=8, + seq_lens=torch.tensor([5, 6], dtype=torch.int32), + block_table=torch.zeros((2, 1), dtype=torch.int32), + slot_mapping=torch.zeros(3, dtype=torch.int32), + use_cascade=False, + common_prefix_len=0, + cu_prefix_query_lens=None, + prefix_kv_lens=None, + suffix_kv_lens=None, + max_dcp_context_kv_len=None, + dcp_context_kv_lens=None, + num_decode_reqs=2, + num_prefill_reqs=0, + num_decode_tokens=3, + num_prefill_tokens=0, + scheduler_metadata=torch.tensor([-1, -1, -1], dtype=torch.int32), + prefix_scheduler_metadata=None, + max_num_splits=4, + causal=True, + sliding_window=None, + mm_prefix_range_tensor=None, + rswa_prefix_lens=None, + rswa_window=None, + rswa_window_tensor=None, + ) + + builder.refresh_meta_for_draft_decodes(metadata) + + assert torch.equal(metadata.scheduler_metadata, expected) + assert torch.equal(builder.scheduler_metadata[:3], expected) + + +def test_refresh_meta_for_draft_decodes_skips_non_fa3_builders(monkeypatch): + builder = object.__new__(flash_attn_module.FlashAttentionMetadataBuilder) + builder.aot_schedule = False + builder.use_full_cuda_graph = True + builder.scheduler_metadata = torch.zeros(4, dtype=torch.int32) + + called = False + + def fake_get_scheduler_metadata(**kwargs): + nonlocal called + called = True + return torch.tensor([1], dtype=torch.int32) + + monkeypatch.setattr(builder, "_get_scheduler_metadata", fake_get_scheduler_metadata) + + metadata = FlashAttentionMetadata( + num_actual_tokens=1, + max_query_len=1, + query_start_loc=torch.tensor([0, 1], dtype=torch.int32), + max_seq_len=1, + seq_lens=torch.tensor([1], dtype=torch.int32), + block_table=torch.zeros((1, 1), dtype=torch.int32), + slot_mapping=torch.zeros(1, dtype=torch.int32), + use_cascade=False, + common_prefix_len=0, + cu_prefix_query_lens=None, + prefix_kv_lens=None, + suffix_kv_lens=None, + max_dcp_context_kv_len=None, + dcp_context_kv_lens=None, + num_decode_reqs=1, + num_prefill_reqs=0, + num_decode_tokens=1, + num_prefill_tokens=0, + scheduler_metadata=torch.tensor([5], dtype=torch.int32), + prefix_scheduler_metadata=None, + max_num_splits=1, + causal=True, + sliding_window=None, + mm_prefix_range_tensor=None, + rswa_prefix_lens=None, + rswa_window=None, + rswa_window_tensor=None, + ) + + builder.refresh_meta_for_draft_decodes(metadata) + + assert not called + assert torch.equal( + metadata.scheduler_metadata, + torch.tensor([5], dtype=torch.int32), + ) diff --git a/vllm/v1/attention/backends/flash_attn.py b/vllm/v1/attention/backends/flash_attn.py index 97f30736fda3..42321da98760 100755 --- a/vllm/v1/attention/backends/flash_attn.py +++ b/vllm/v1/attention/backends/flash_attn.py @@ -370,6 +370,57 @@ def get_cudagraph_support( ) -> AttentionCGSupport: return cls._cudagraph_support + def _get_scheduler_metadata( + self, + *, + aot_schedule: bool, + batch_size: int, + cu_query_lens: torch.Tensor, + max_query_len: int, + seqlens: torch.Tensor, + max_seq_len: int, + causal: bool | torch.Tensor, + max_num_splits: int, + ) -> torch.Tensor | None: + if not aot_schedule: + return None + + cache_dtype = self.cache_config.cache_dtype + if is_quantized_kv_cache(cache_dtype): + qkv_dtype = current_platform.fp8_dtype() + else: + qkv_dtype = self.kv_cache_dtype + return get_scheduler_metadata( + batch_size=batch_size, + max_seqlen_q=max_query_len, + max_seqlen_k=max_seq_len, + num_heads_q=self.num_heads_q * self.dcp_world_size, + num_heads_kv=self.num_heads_kv, + headdim=self.headdim, + cache_seqlens=seqlens, + qkv_dtype=qkv_dtype, + cu_seqlens_q=cu_query_lens, + page_size=self.block_size, + causal=causal, + window_size=_maybe_symmetrize_window(self.aot_sliding_window, causal), + num_splits=max_num_splits, + ) + + def _store_scheduler_metadata( + self, scheduler_metadata: torch.Tensor | None + ) -> torch.Tensor | None: + if self.use_full_cuda_graph and scheduler_metadata is not None: + n = scheduler_metadata.shape[0] + assert self.scheduler_metadata is not None + self.scheduler_metadata[:n] = scheduler_metadata + # NOTE(woosuk): We should zero out the rest of the scheduler + # metadata to guarantee the correctness. Otherwise, some thread + # blocks may use the invalid scheduler metadata and overwrite the + # output buffer. + self.scheduler_metadata[n:] = 0 + return self.scheduler_metadata[:n] + return scheduler_metadata + def __init__( self, kv_cache_spec: AttentionSpec, @@ -405,6 +456,14 @@ def __init__( self.dcp_world_size = 1 self.dcp_rank = 0 + # Fused draft decode reuses the captured metadata object across draft + # steps. For DCP, build-time host-side decisions such as + # skip_dcp_context_attention() can change the metadata shape/control + # path (for example max_dcp_context_kv_len), and those Python-side + # fields are not refreshed in-place between graph replays. Keep the + # fused path disabled until DCP gets a full replay-safe refresh model. + self.supports_fused_decode_graph = self.dcp_world_size == 1 + self.cp_kv_cache_interleave_size = ( self.parallel_config.cp_kv_cache_interleave_size ) @@ -533,34 +592,6 @@ def build( if envs.VLLM_BATCH_INVARIANT: max_num_splits = 1 - def schedule( - batch_size, cu_query_lens, max_query_len, seqlens, max_seq_len, causal - ): - cache_dtype = self.cache_config.cache_dtype - if is_quantized_kv_cache(cache_dtype): - qkv_dtype = current_platform.fp8_dtype() - else: - qkv_dtype = self.kv_cache_dtype - if aot_schedule: - return get_scheduler_metadata( - batch_size=batch_size, - max_seqlen_q=max_query_len, - max_seqlen_k=max_seq_len, - num_heads_q=self.num_heads_q * self.dcp_world_size, - num_heads_kv=self.num_heads_kv, - headdim=self.headdim, - cache_seqlens=seqlens, - qkv_dtype=qkv_dtype, - cu_seqlens_q=cu_query_lens, - page_size=self.block_size, - causal=causal, - window_size=_maybe_symmetrize_window( - self.aot_sliding_window, causal - ), - num_splits=max_num_splits, - ) - return None - use_cascade = common_prefix_len > 0 max_dcp_context_kv_len = 0 dcp_context_kv_lens = None @@ -627,13 +658,15 @@ def schedule( (max_seq_len + num_partitions - 1) // num_partitions ) * self.cp_kv_cache_interleave_size - scheduler_metadata = schedule( + scheduler_metadata = self._get_scheduler_metadata( + aot_schedule=aot_schedule, batch_size=num_reqs, cu_query_lens=query_start_loc, max_query_len=max_query_len, seqlens=dcp_context_kv_lens, max_seq_len=max_dcp_context_kv_len, causal=False, + max_num_splits=max_num_splits, ) elif use_cascade: cu_prefix_query_lens = torch.tensor( @@ -644,41 +677,38 @@ def schedule( ) # Use GPU tensor directly - no CPU sync needed suffix_kv_lens = seq_lens[:num_reqs] - common_prefix_len - prefix_scheduler_metadata = schedule( + prefix_scheduler_metadata = self._get_scheduler_metadata( + aot_schedule=aot_schedule, batch_size=1, cu_query_lens=cu_prefix_query_lens, max_query_len=num_actual_tokens, seqlens=prefix_kv_lens, max_seq_len=common_prefix_len, causal=False, + max_num_splits=max_num_splits, ) - scheduler_metadata = schedule( + scheduler_metadata = self._get_scheduler_metadata( + aot_schedule=aot_schedule, batch_size=num_reqs, cu_query_lens=query_start_loc, max_query_len=max_query_len, seqlens=suffix_kv_lens, max_seq_len=max_seq_len - common_prefix_len, causal=True, + max_num_splits=max_num_splits, ) else: - scheduler_metadata = schedule( + scheduler_metadata = self._get_scheduler_metadata( + aot_schedule=aot_schedule, batch_size=num_reqs, cu_query_lens=query_start_loc, max_query_len=max_query_len, seqlens=seq_lens, max_seq_len=max_seq_len, causal=causal, + max_num_splits=max_num_splits, ) - # For FA3 + full cudagraph - if self.use_full_cuda_graph and scheduler_metadata is not None: - n = scheduler_metadata.shape[0] - self.scheduler_metadata[:n] = scheduler_metadata - # NOTE(woosuk): We should zero out the rest of the scheduler - # metadata to guarantee the correctness. Otherwise, some thread - # blocks may use the invalid scheduler metadata and overwrite the - # output buffer. - self.scheduler_metadata[n:] = 0 - scheduler_metadata = self.scheduler_metadata[:n] + scheduler_metadata = self._store_scheduler_metadata(scheduler_metadata) if isinstance(causal, torch.Tensor) and causal.dtype != torch.int32: causal = causal.to(torch.int32) @@ -772,6 +802,51 @@ def update_block_table( new_metadata.slot_mapping = slot_mapping return new_metadata + def refresh_meta_for_draft_decodes(self, metadata: FlashAttentionMetadata) -> None: + if not self.aot_schedule: + return + + num_reqs = metadata.num_decode_reqs or metadata.seq_lens.shape[0] + + if metadata.use_cascade: + assert metadata.cu_prefix_query_lens is not None + assert metadata.prefix_kv_lens is not None + assert metadata.suffix_kv_lens is not None + prefix_scheduler_metadata = self._get_scheduler_metadata( + aot_schedule=True, + batch_size=1, + cu_query_lens=metadata.cu_prefix_query_lens, + max_query_len=metadata.num_actual_tokens, + seqlens=metadata.prefix_kv_lens, + max_seq_len=metadata.common_prefix_len, + causal=False, + max_num_splits=metadata.max_num_splits, + ) + metadata.prefix_scheduler_metadata = prefix_scheduler_metadata + scheduler_metadata = self._get_scheduler_metadata( + aot_schedule=True, + batch_size=num_reqs, + cu_query_lens=metadata.query_start_loc, + max_query_len=metadata.max_query_len, + seqlens=metadata.suffix_kv_lens, + max_seq_len=metadata.max_seq_len - metadata.common_prefix_len, + causal=True, + max_num_splits=metadata.max_num_splits, + ) + else: + scheduler_metadata = self._get_scheduler_metadata( + aot_schedule=True, + batch_size=num_reqs, + cu_query_lens=metadata.query_start_loc, + max_query_len=metadata.max_query_len, + seqlens=metadata.seq_lens, + max_seq_len=metadata.max_seq_len, + causal=metadata.causal, + max_num_splits=metadata.max_num_splits, + ) + + metadata.scheduler_metadata = self._store_scheduler_metadata(scheduler_metadata) + def use_cascade_attention(self, *args, **kwargs) -> bool: return use_cascade_attention(*args, **kwargs) From 94ee31420c40ca54ff5a5773d5608618f96c1751 Mon Sep 17 00:00:00 2001 From: Yizhou Liu Date: Tue, 28 Jul 2026 03:31:46 +0000 Subject: [PATCH 03/11] feat(spec): refine fused draft decode dispatch Signed-off-by: Yizhou Liu --- .../test_gpu_autoregressive_speculator.py | 14 +-- vllm/config/speculative.py | 4 - .../spec_decode/autoregressive/speculator.py | 103 ++++++++++++------ 3 files changed, 74 insertions(+), 47 deletions(-) diff --git a/tests/v1/worker/test_gpu_autoregressive_speculator.py b/tests/v1/worker/test_gpu_autoregressive_speculator.py index bff851592395..ad7fecb1af91 100644 --- a/tests/v1/worker/test_gpu_autoregressive_speculator.py +++ b/tests/v1/worker/test_gpu_autoregressive_speculator.py @@ -286,20 +286,21 @@ def test_run_model_reuses_tensor_return_for_mtp(monkeypatch): @pytest.mark.parametrize( ( + "method_name", "cg_mode", - "use_fused_decode_graph", "expected_eager_calls", "expected_graph_replays", ), [ - (CUDAGraphMode.NONE, True, 3, 0), - (CUDAGraphMode.FULL, True, 0, 1), - (CUDAGraphMode.FULL, False, 0, 3), + ("_multi_step_decode", CUDAGraphMode.NONE, 3, 0), + ("_multi_step_decode", CUDAGraphMode.FULL, 0, 3), + ("_fused_multi_step_decode", CUDAGraphMode.NONE, 3, 0), + ("_fused_multi_step_decode", CUDAGraphMode.FULL, 0, 1), ], ) def test_multi_step_decode_replays_captured_graph_as_expected( + method_name, cg_mode, - use_fused_decode_graph, expected_eager_calls, expected_graph_replays, ): @@ -311,7 +312,6 @@ def test_multi_step_decode_replays_captured_graph_as_expected( query_start_loc=torch.arange(3), ) speculator.idx_mapping = torch.arange(2) - speculator.use_fused_decode_graph = use_fused_decode_graph generate_draft = Mock() speculator._generate_draft = generate_draft run_fullgraph = Mock() @@ -322,7 +322,7 @@ def test_multi_step_decode_replays_captured_graph_as_expected( num_reqs=2, ) - speculator._multi_step_decode( + getattr(speculator, method_name)( num_reqs=2, skip_attn=True, batch_desc=batch_desc, diff --git a/vllm/config/speculative.py b/vllm/config/speculative.py index 1a817e759adc..fa8c282007b2 100644 --- a/vllm/config/speculative.py +++ b/vllm/config/speculative.py @@ -138,9 +138,6 @@ class SpeculativeConfig: will use the default version.""" # Advanced control - enable_fused_decode_graph: bool = True - """Fuse autoregressive draft decode steps into one CUDA graph when the - attention backends support it. Unsupported backends use per-step graphs.""" disable_padded_drafter_batch: bool = False """Disable input padding for speculative decoding. If set to True, speculative input batches can contain sequences of different lengths, @@ -333,7 +330,6 @@ def compute_hash(self) -> str: # Convert to tuple to make it hashable factors.append(tuple(layer_ids)) - factors.append(self.enable_fused_decode_graph) hash_str = safe_hash(str(factors).encode(), usedforsecurity=False).hexdigest() return hash_str diff --git a/vllm/v1/worker/gpu/spec_decode/autoregressive/speculator.py b/vllm/v1/worker/gpu/spec_decode/autoregressive/speculator.py index f96310d49c38..9ea30f65a395 100644 --- a/vllm/v1/worker/gpu/spec_decode/autoregressive/speculator.py +++ b/vllm/v1/worker/gpu/spec_decode/autoregressive/speculator.py @@ -99,10 +99,7 @@ def set_attn( self._configure_fused_decode_graph() def _configure_fused_decode_graph(self) -> None: - if ( - not self.speculative_config.enable_fused_decode_graph - or self.num_speculative_steps == 1 - ): + if self.num_speculative_steps == 1: self.use_fused_decode_graph = False return @@ -183,7 +180,7 @@ def capture(self) -> None: # Capture either the fused decode loop or one decode step per graph. assert self.decode_cudagraph_manager is not None decode_fn = ( - self._run_draft_decode_loop + self._generate_fused_drafts if self.use_fused_decode_graph else self._generate_draft ) @@ -340,7 +337,12 @@ def propose( self.on_multi_step_decode_begin(num_reqs) # Generate the remaining num_speculative_steps - 1 draft tokens. - self._multi_step_decode( + decode_fn = ( + self._fused_multi_step_decode + if self.use_fused_decode_graph + else self._multi_step_decode + ) + decode_fn( num_reqs, dummy_run and skip_attn_for_dummy_run, decode_batch_desc, @@ -462,27 +464,6 @@ def _multi_step_decode( query_start_loc = self.input_buffers.query_start_loc[: num_reqs + 1] idx_mapping = self.idx_mapping[:num_reqs] - if batch_desc.cg_mode == CUDAGraphMode.FULL and self.use_fused_decode_graph: - if not skip_attn: - self.block_tables.compute_slot_mappings( - idx_mapping, - query_start_loc, - positions, - batch_desc.num_tokens, - ) - # Continuous draft decode replays usually consume only device metadata, - # so the host-side vars should have no effect and therefore pinned. - self._build_draft_attn_metadata( - num_reqs=num_reqs, - num_reqs_padded=batch_desc.num_reqs or num_reqs, - num_tokens_padded=batch_desc.num_tokens, - seq_lens_cpu_upper_bound=seq_lens_cpu_upper_bound, - step=1, - ) - assert self.decode_cudagraph_manager is not None - self.decode_cudagraph_manager.run_fullgraph(batch_desc) - return - attn_metadata = None slot_mappings_by_layer = None for step in range(1, self.num_speculative_steps): @@ -521,7 +502,54 @@ def _multi_step_decode( cudagraph_runtime_mode=batch_desc.cg_mode, ) - def _run_draft_decode_loop( + def _fused_multi_step_decode( + self, + num_reqs: int, + skip_attn: bool, + batch_desc: BatchExecutionDescriptor, + num_tokens_across_dp: torch.Tensor | None, + seq_lens_cpu_upper_bound: torch.Tensor, + ) -> None: + positions = self.input_buffers.positions[:num_reqs] + query_start_loc = self.input_buffers.query_start_loc[: num_reqs + 1] + idx_mapping = self.idx_mapping[:num_reqs] + + attn_metadata = None + slot_mappings_by_layer = None + if not skip_attn: + slot_mappings = self.block_tables.compute_slot_mappings( + idx_mapping, + query_start_loc, + positions, + batch_desc.num_tokens, + ) + if batch_desc.cg_mode != CUDAGraphMode.FULL: + slot_mappings_by_layer = build_slot_mappings_by_layer( + slot_mappings, self.kv_cache_config + ) + attn_metadata = self._build_draft_attn_metadata( + num_reqs=num_reqs, + num_reqs_padded=batch_desc.num_reqs or num_reqs, + num_tokens_padded=batch_desc.num_tokens, + seq_lens_cpu_upper_bound=seq_lens_cpu_upper_bound, + step=1, + ) + + if batch_desc.cg_mode == CUDAGraphMode.FULL: + assert self.decode_cudagraph_manager is not None + self.decode_cudagraph_manager.run_fullgraph(batch_desc) + return + + self._generate_fused_drafts( + num_reqs, + batch_desc.num_tokens, + attn_metadata, + slot_mappings_by_layer, + num_tokens_across_dp, + batch_desc.cg_mode, + ) + + def _generate_fused_drafts( self, num_reqs: int, num_tokens_padded: int, @@ -549,14 +577,17 @@ def _run_draft_decode_loop( num_tokens_across_dp, cudagraph_runtime_mode, ) - if step < self.num_speculative_steps - 1 and attn_metadata is not None: - if self.advance_draft_positions: - self.block_tables.compute_slot_mappings( - idx_mapping, - query_start_loc, - positions, - num_tokens_padded, - ) + if ( + step < self.num_speculative_steps - 1 + and attn_metadata is not None + and self.advance_draft_positions + ): + self.block_tables.compute_slot_mappings( + idx_mapping, + query_start_loc, + positions, + num_tokens_padded, + ) for attn_group in attn_groups: attn_group.refresh_meta_for_draft_decodes(attn_metadata) From e10375fa69012afe0313fff15eab63fc1abf3dcd Mon Sep 17 00:00:00 2001 From: Yizhou Liu Date: Tue, 28 Jul 2026 04:42:32 +0000 Subject: [PATCH 04/11] feat(spec): restrict FA3 draft metadata refresh Signed-off-by: Yizhou Liu --- vllm/v1/attention/backends/flash_attn.py | 49 +++++++----------------- 1 file changed, 13 insertions(+), 36 deletions(-) diff --git a/vllm/v1/attention/backends/flash_attn.py b/vllm/v1/attention/backends/flash_attn.py index 42321da98760..24f003a61158 100755 --- a/vllm/v1/attention/backends/flash_attn.py +++ b/vllm/v1/attention/backends/flash_attn.py @@ -808,42 +808,19 @@ def refresh_meta_for_draft_decodes(self, metadata: FlashAttentionMetadata) -> No num_reqs = metadata.num_decode_reqs or metadata.seq_lens.shape[0] - if metadata.use_cascade: - assert metadata.cu_prefix_query_lens is not None - assert metadata.prefix_kv_lens is not None - assert metadata.suffix_kv_lens is not None - prefix_scheduler_metadata = self._get_scheduler_metadata( - aot_schedule=True, - batch_size=1, - cu_query_lens=metadata.cu_prefix_query_lens, - max_query_len=metadata.num_actual_tokens, - seqlens=metadata.prefix_kv_lens, - max_seq_len=metadata.common_prefix_len, - causal=False, - max_num_splits=metadata.max_num_splits, - ) - metadata.prefix_scheduler_metadata = prefix_scheduler_metadata - scheduler_metadata = self._get_scheduler_metadata( - aot_schedule=True, - batch_size=num_reqs, - cu_query_lens=metadata.query_start_loc, - max_query_len=metadata.max_query_len, - seqlens=metadata.suffix_kv_lens, - max_seq_len=metadata.max_seq_len - metadata.common_prefix_len, - causal=True, - max_num_splits=metadata.max_num_splits, - ) - else: - scheduler_metadata = self._get_scheduler_metadata( - aot_schedule=True, - batch_size=num_reqs, - cu_query_lens=metadata.query_start_loc, - max_query_len=metadata.max_query_len, - seqlens=metadata.seq_lens, - max_seq_len=metadata.max_seq_len, - causal=metadata.causal, - max_num_splits=metadata.max_num_splits, - ) + assert self.dcp_world_size == 1 + assert not metadata.use_cascade + + scheduler_metadata = self._get_scheduler_metadata( + aot_schedule=True, + batch_size=num_reqs, + cu_query_lens=metadata.query_start_loc, + max_query_len=metadata.max_query_len, + seqlens=metadata.seq_lens, + max_seq_len=metadata.max_seq_len, + causal=metadata.causal, + max_num_splits=metadata.max_num_splits, + ) metadata.scheduler_metadata = self._store_scheduler_metadata(scheduler_metadata) From ee724b18960dc4247a8e5fcb15ef8091223a8d59 Mon Sep 17 00:00:00 2001 From: Yizhou Liu Date: Tue, 4 Aug 2026 09:41:36 +0800 Subject: [PATCH 05/11] feat(spec): refine fused draft metadata update contract Signed-off-by: Yizhou Liu --- .../test_gpu_autoregressive_speculator.py | 8 +++---- vllm/v1/attention/backend.py | 11 ++++----- vllm/v1/attention/backends/flash_attn.py | 4 ++-- vllm/v1/attention/backends/mla/sparse_swa.py | 4 ++-- .../spec_decode/autoregressive/speculator.py | 24 +++++++++++-------- vllm/v1/worker/utils.py | 8 +++---- 6 files changed, 31 insertions(+), 28 deletions(-) diff --git a/tests/v1/worker/test_gpu_autoregressive_speculator.py b/tests/v1/worker/test_gpu_autoregressive_speculator.py index ad7fecb1af91..11d2f00e2e3a 100644 --- a/tests/v1/worker/test_gpu_autoregressive_speculator.py +++ b/tests/v1/worker/test_gpu_autoregressive_speculator.py @@ -334,7 +334,7 @@ def test_multi_step_decode_replays_captured_graph_as_expected( assert run_fullgraph.call_count == expected_graph_replays -def test_refresh_meta_for_draft_decodes_updates_fa3_scheduler_metadata( +def test_update_draft_decode_metadata_updates_fa3_scheduler_metadata( monkeypatch, ): builder = object.__new__(flash_attn_module.FlashAttentionMetadataBuilder) @@ -389,13 +389,13 @@ def fake_get_scheduler_metadata(**kwargs): rswa_window_tensor=None, ) - builder.refresh_meta_for_draft_decodes(metadata) + builder.update_draft_decode_metadata(metadata) assert torch.equal(metadata.scheduler_metadata, expected) assert torch.equal(builder.scheduler_metadata[:3], expected) -def test_refresh_meta_for_draft_decodes_skips_non_fa3_builders(monkeypatch): +def test_update_draft_decode_metadata_skips_non_fa3_builders(monkeypatch): builder = object.__new__(flash_attn_module.FlashAttentionMetadataBuilder) builder.aot_schedule = False builder.use_full_cuda_graph = True @@ -440,7 +440,7 @@ def fake_get_scheduler_metadata(**kwargs): rswa_window_tensor=None, ) - builder.refresh_meta_for_draft_decodes(metadata) + builder.update_draft_decode_metadata(metadata) assert not called assert torch.equal( diff --git a/vllm/v1/attention/backend.py b/vllm/v1/attention/backend.py index 25a7135c8aef..a6f30b60030a 100644 --- a/vllm/v1/attention/backend.py +++ b/vllm/v1/attention/backend.py @@ -633,10 +633,9 @@ class AttentionMetadataBuilder(ABC, Generic[M]): supports_update_block_table: bool = False # Whether the builder constructor requires the block-table width. requires_block_table_width: ClassVar[bool] = False - # Does this backend support capture mutiple draft decode steps into one - # CUDA Graph (default: no), which requires no step-dependent metadata - # or can refresh metadata over steps. - supports_fused_decode_graph: bool = False + # Whether all step-dependent draft decode metadata can be updated in place, + # allowing one metadata build to be reused across autoregressive draft steps. + supports_draft_decode_metadata_update: bool = False @abstractmethod def __init__( @@ -761,8 +760,8 @@ def build_for_drafting( fast_build=True, ) - def refresh_meta_for_draft_decodes(self, metadata: M) -> None: - pass + def update_draft_decode_metadata(self, metadata: M) -> None: + raise NotImplementedError def use_cascade_attention( self, diff --git a/vllm/v1/attention/backends/flash_attn.py b/vllm/v1/attention/backends/flash_attn.py index 24f003a61158..222b2c64b5d0 100755 --- a/vllm/v1/attention/backends/flash_attn.py +++ b/vllm/v1/attention/backends/flash_attn.py @@ -462,7 +462,7 @@ def __init__( # path (for example max_dcp_context_kv_len), and those Python-side # fields are not refreshed in-place between graph replays. Keep the # fused path disabled until DCP gets a full replay-safe refresh model. - self.supports_fused_decode_graph = self.dcp_world_size == 1 + self.supports_draft_decode_metadata_update = self.dcp_world_size == 1 self.cp_kv_cache_interleave_size = ( self.parallel_config.cp_kv_cache_interleave_size @@ -802,7 +802,7 @@ def update_block_table( new_metadata.slot_mapping = slot_mapping return new_metadata - def refresh_meta_for_draft_decodes(self, metadata: FlashAttentionMetadata) -> None: + def update_draft_decode_metadata(self, metadata: FlashAttentionMetadata) -> None: if not self.aot_schedule: return diff --git a/vllm/v1/attention/backends/mla/sparse_swa.py b/vllm/v1/attention/backends/mla/sparse_swa.py index d7abcc07fcf0..cbd4fc467409 100644 --- a/vllm/v1/attention/backends/mla/sparse_swa.py +++ b/vllm/v1/attention/backends/mla/sparse_swa.py @@ -400,7 +400,7 @@ class DeepseekSparseSWAMetadataBuilder(AttentionMetadataBuilder): reorder_batch_threshold: int | None = None _cudagraph_support: ClassVar[AttentionCGSupport] = AttentionCGSupport.UNIFORM_BATCH - supports_fused_decode_graph = True + supports_draft_decode_metadata_update = True def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) @@ -650,7 +650,7 @@ def build( **deepseek_v4_fields, # type: ignore[arg-type] ) - def refresh_meta_for_draft_decodes( + def update_draft_decode_metadata( self, metadata: DeepseekSparseSWAMetadata, ) -> None: diff --git a/vllm/v1/worker/gpu/spec_decode/autoregressive/speculator.py b/vllm/v1/worker/gpu/spec_decode/autoregressive/speculator.py index 9ea30f65a395..8dcde41d5771 100644 --- a/vllm/v1/worker/gpu/spec_decode/autoregressive/speculator.py +++ b/vllm/v1/worker/gpu/spec_decode/autoregressive/speculator.py @@ -45,7 +45,7 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device): self.prefill_cudagraph_manager: SpeculatorCudaGraphManager | None = None self.decode_cudagraph_manager: SpeculatorCudaGraphManager | None = None - self.use_fused_decode_graph = False + self.use_fused_multi_step_decode = False def load_model(self, target_model: nn.Module) -> None: super().load_model(target_model) @@ -96,11 +96,15 @@ def set_attn( target_input_buffers, target_attn_groups, ) - self._configure_fused_decode_graph() + self._configure_fused_multi_step_decode() - def _configure_fused_decode_graph(self) -> None: + def _configure_fused_multi_step_decode(self) -> None: if self.num_speculative_steps == 1: - self.use_fused_decode_graph = False + self.use_fused_multi_step_decode = False + return + + if not self.advance_draft_positions: + self.use_fused_multi_step_decode = True return unsupported_backends = sorted( @@ -108,13 +112,13 @@ def _configure_fused_decode_graph(self) -> None: attn_group.backend.get_name() for attn_groups in self.attn_groups for attn_group in attn_groups - if not attn_group.supports_fused_decode_graph + if not attn_group.supports_draft_decode_metadata_update } ) - self.use_fused_decode_graph = not unsupported_backends + self.use_fused_multi_step_decode = not unsupported_backends if unsupported_backends: logger.info_once( - "Fused draft decode graph is not supported by attention " + "Fused multi-step draft decode is not supported by attention " "backend(s) %s; falling back to rebuilding attention metadata " "between draft steps.", ", ".join(unsupported_backends), @@ -181,7 +185,7 @@ def capture(self) -> None: assert self.decode_cudagraph_manager is not None decode_fn = ( self._generate_fused_drafts - if self.use_fused_decode_graph + if self.use_fused_multi_step_decode else self._generate_draft ) self.decode_cudagraph_manager.capture( @@ -339,7 +343,7 @@ def propose( # Generate the remaining num_speculative_steps - 1 draft tokens. decode_fn = ( self._fused_multi_step_decode - if self.use_fused_decode_graph + if self.use_fused_multi_step_decode else self._multi_step_decode ) decode_fn( @@ -589,7 +593,7 @@ def _generate_fused_drafts( num_tokens_padded, ) for attn_group in attn_groups: - attn_group.refresh_meta_for_draft_decodes(attn_metadata) + attn_group.update_draft_decode_metadata(attn_metadata) def _generate_draft( self, diff --git a/vllm/v1/worker/utils.py b/vllm/v1/worker/utils.py index f0b12cda238c..55f93860ad79 100644 --- a/vllm/v1/worker/utils.py +++ b/vllm/v1/worker/utils.py @@ -300,15 +300,15 @@ def get_metadata_builder(self, ubatch_id: int = 0) -> AttentionMetadataBuilder: return self.metadata_builders[ubatch_id] @property - def supports_fused_decode_graph(self) -> bool: - return self.get_metadata_builder().supports_fused_decode_graph + def supports_draft_decode_metadata_update(self) -> bool: + return self.get_metadata_builder().supports_draft_decode_metadata_update - def refresh_meta_for_draft_decodes( + def update_draft_decode_metadata( self, attn_metadata: Mapping[str, Any], ) -> None: metadata = attn_metadata[self.layer_names[0]] - self.get_metadata_builder().refresh_meta_for_draft_decodes(metadata) + self.get_metadata_builder().update_draft_decode_metadata(metadata) def select_common_block_size( From 1e19f42a9a3c54b081ffbabfb6cdd424b5a70e23 Mon Sep 17 00:00:00 2001 From: Yizhou Liu Date: Tue, 4 Aug 2026 09:42:22 +0800 Subject: [PATCH 06/11] fix(spec): disable fused draft metadata updates on ROCm Signed-off-by: Yizhou Liu --- vllm/models/deepseek_v4/amd/rocm.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/vllm/models/deepseek_v4/amd/rocm.py b/vllm/models/deepseek_v4/amd/rocm.py index 7c232fe9e528..9b38409d1fde 100644 --- a/vllm/models/deepseek_v4/amd/rocm.py +++ b/vllm/models/deepseek_v4/amd/rocm.py @@ -372,6 +372,10 @@ def build( class DeepseekV4ROCMAiterSparseSWAMetadataBuilder(DeepseekSparseSWAMetadataBuilder): + # Keep fused multi-step decode disabled until update_draft_decode_metadata() + # also refreshes the ROCm-specific ragged SWA indices and indptrs. + supports_draft_decode_metadata_update = False + def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) max_tokens = self.vllm_config.scheduler_config.max_num_batched_tokens From 1b8efd0065fc20e3dbe40571b6404c05adf588e6 Mon Sep 17 00:00:00 2001 From: Yizhou Liu Date: Thu, 6 Aug 2026 20:50:36 +0800 Subject: [PATCH 07/11] fix(spec): invalidate FlashInfer sparse indices between draft steps Signed-off-by: Yizhou Liu --- vllm/v1/attention/backends/mla/sparse_swa.py | 1 + 1 file changed, 1 insertion(+) diff --git a/vllm/v1/attention/backends/mla/sparse_swa.py b/vllm/v1/attention/backends/mla/sparse_swa.py index cbd4fc467409..c11698335a0b 100644 --- a/vllm/v1/attention/backends/mla/sparse_swa.py +++ b/vllm/v1/attention/backends/mla/sparse_swa.py @@ -682,6 +682,7 @@ def update_draft_decode_metadata( metadata.tile_sched_swaonly = tile_sched[_LAYER_TYPE_SWAONLY] metadata.tile_sched_c4a = tile_sched[_LAYER_TYPE_C4A] metadata.tile_sched_c128a = tile_sched[_LAYER_TYPE_C128A] + metadata.flashinfer_sparse_index_cache.clear() def build_tile_scheduler( self, num_decode_tokens: int From 8236b3f68ed3b39e96f0da981f61dca86bd55b75 Mon Sep 17 00:00:00 2001 From: Yizhou Liu Date: Thu, 6 Aug 2026 21:21:09 +0800 Subject: [PATCH 08/11] feat(spec): enable fused draft decode for Triton backends Signed-off-by: Yizhou Liu --- vllm/v1/attention/backends/mla/triton_mla.py | 5 +++++ vllm/v1/attention/backends/triton_attn.py | 5 +++++ 2 files changed, 10 insertions(+) diff --git a/vllm/v1/attention/backends/mla/triton_mla.py b/vllm/v1/attention/backends/mla/triton_mla.py index 71a38e95d0ad..e55076700fd7 100644 --- a/vllm/v1/attention/backends/mla/triton_mla.py +++ b/vllm/v1/attention/backends/mla/triton_mla.py @@ -57,6 +57,8 @@ class TritonMLAMetadataBuilder(MLACommonMetadataBuilder[MLACommonMetadata]): def __init__(self, kv_cache_spec, layer_names, vllm_config, device): super().__init__(kv_cache_spec, layer_names, vllm_config, device) + # DCP local sequence lengths are not advanced between draft steps. + self.supports_draft_decode_metadata_update = self.dcp_world_size == 1 # Only the non-causal DSpark draft group serves multi-token blocks via # the decode path; raise its reorder threshold to the spec block length # so full-cudagraph capture admits it. Causal usage stays single-token. @@ -64,6 +66,9 @@ def __init__(self, kv_cache_spec, layer_names, vllm_config, device): self._init_reorder_batch_threshold(1, supports_spec_as_decode=True) self._reserve_attn_logits_workspace() + def update_draft_decode_metadata(self, _metadata: MLACommonMetadata) -> None: + pass + def _reserve_attn_logits_workspace(self) -> None: """Pre-size the shared workspace for the decode split-KV attn logits. diff --git a/vllm/v1/attention/backends/triton_attn.py b/vllm/v1/attention/backends/triton_attn.py index 3f9b9ed0c144..7814a8c0f8de 100644 --- a/vllm/v1/attention/backends/triton_attn.py +++ b/vllm/v1/attention/backends/triton_attn.py @@ -100,6 +100,8 @@ class TritonAttentionMetadata: class TritonAttentionMetadataBuilder(AttentionMetadataBuilder[TritonAttentionMetadata]): _cudagraph_support: ClassVar[AttentionCGSupport] = AttentionCGSupport.ALWAYS + # Step-dependent fields reference persistent input buffers directly. + supports_draft_decode_metadata_update = True def __init__( self, @@ -267,6 +269,9 @@ def build( return attn_metadata + def update_draft_decode_metadata(self, _metadata: TritonAttentionMetadata) -> None: + pass + class TritonAttentionBackend(AttentionBackend): supported_dtypes: ClassVar[list[torch.dtype]] = [ From 9117018d2990a06c5f92445e207d566db3dad04e Mon Sep 17 00:00:00 2001 From: Yizhou Liu Date: Thu, 6 Aug 2026 21:34:42 +0800 Subject: [PATCH 09/11] docs(spec): document fused draft graph metadata contract Signed-off-by: Yizhou Liu --- docs/design/model_runner_v2.md | 8 ++++++++ vllm/v1/attention/backend.py | 7 +++++++ 2 files changed, 15 insertions(+) diff --git a/docs/design/model_runner_v2.md b/docs/design/model_runner_v2.md index fb40d51ee7b7..cabc3526b19e 100644 --- a/docs/design/model_runner_v2.md +++ b/docs/design/model_runner_v2.md @@ -191,6 +191,14 @@ V1's CUDA graph handling is implicit and hard to reason about. MRV2 uses a `CUDA This makes graph lifecycle and execution mode decisions more understandable and easier to extend. Example: MRV2 can capture multiple draft-model forward passes into one CUDA graph. +### Fused Multi-Step Draft Decoding + +Autoregressive speculative decoding executes several dependent draft steps per scheduler step. In the fused path, MRV2 captures all post-prefill draft steps in one full CUDA graph instead of replaying a separate graph for each draft token. Attention metadata is built once before the loop, while common step-dependent tensors keep stable addresses and are updated in place between draft steps. + +Some attention backends also materialize derived state, such as scheduler metadata or sparse indices. Before opting into the fused path, these backends must implement `AttentionMetadataBuilder.update_draft_decode_metadata()` to update or invalidate that state after the draft inputs advance. The hook runs during CUDA graph capture, so only the GPU operations it issues are recorded and executed during replay; its Python body is not run again. Implementations must therefore use capture-safe operations and keep all replayed tensor state in persistent storage. + +For draft models that advance positions, the fused path is enabled only when every draft attention group declares `supports_draft_decode_metadata_update`. Otherwise, MRV2 falls back to rebuilding attention metadata between draft steps. Draft models that keep positions fixed do not require this update. Before enabling a backend, developers must audit all derived metadata, including state inherited from parent builders or owned by auxiliary attention backends. + ## Development Philosophy MRV2 changes should meet a higher code quality bar. As feature gaps with V1 are filled, features should be reconsidered from first principles in the MRV2 design context instead of quickly porting V1 behavior. diff --git a/vllm/v1/attention/backend.py b/vllm/v1/attention/backend.py index a6f30b60030a..32e932e3b898 100644 --- a/vllm/v1/attention/backend.py +++ b/vllm/v1/attention/backend.py @@ -761,6 +761,13 @@ def build_for_drafting( ) def update_draft_decode_metadata(self, metadata: M) -> None: + """Update step-dependent draft decode metadata in place. + + The fused draft loop may call this method during full CUDA graph + capture. CUDA graph replay does not run this Python method, so + implementations must emit capture-safe operations and keep replayed + tensor state in persistent storage. + """ raise NotImplementedError def use_cascade_attention( From 0fe9ff4ee61245aba7a66c63771637eb3d513669 Mon Sep 17 00:00:00 2001 From: Yizhou Liu Date: Sat, 8 Aug 2026 19:34:16 +0800 Subject: [PATCH 10/11] fix(spec): reset index mapping before graph capture This became necessary after #48892 made padded idx_mapping entries persist as -1 for all draft sampling modes. Signed-off-by: Yizhou Liu --- vllm/v1/worker/gpu/spec_decode/autoregressive/speculator.py | 1 + 1 file changed, 1 insertion(+) diff --git a/vllm/v1/worker/gpu/spec_decode/autoregressive/speculator.py b/vllm/v1/worker/gpu/spec_decode/autoregressive/speculator.py index 8dcde41d5771..7d357533e27e 100644 --- a/vllm/v1/worker/gpu/spec_decode/autoregressive/speculator.py +++ b/vllm/v1/worker/gpu/spec_decode/autoregressive/speculator.py @@ -152,6 +152,7 @@ def capture(self) -> None: # Reset indices to zeros to prevent stale values from prior # dummy runs to cause out-of-bounds indexing during capture. self.last_token_indices.zero_() + self.idx_mapping.zero_() # Capture the prefill routine (model forward + compute_logits + # sample). From fc8ceaf0fc498b229913d02ada293a548521f8df Mon Sep 17 00:00:00 2001 From: Yizhou Liu Date: Tue, 11 Aug 2026 12:50:32 +0800 Subject: [PATCH 11/11] fix(spec): refresh available FA3 scheduler metadata Signed-off-by: Yizhou Liu --- .../worker/test_gpu_autoregressive_speculator.py | 15 ++++++--------- vllm/v1/attention/backends/flash_attn.py | 2 +- 2 files changed, 7 insertions(+), 10 deletions(-) diff --git a/tests/v1/worker/test_gpu_autoregressive_speculator.py b/tests/v1/worker/test_gpu_autoregressive_speculator.py index 11d2f00e2e3a..0d8c2943b0d1 100644 --- a/tests/v1/worker/test_gpu_autoregressive_speculator.py +++ b/tests/v1/worker/test_gpu_autoregressive_speculator.py @@ -383,7 +383,7 @@ def fake_get_scheduler_metadata(**kwargs): max_num_splits=4, causal=True, sliding_window=None, - mm_prefix_range_tensor=None, + mm_prefix_query_range_tensor=None, rswa_prefix_lens=None, rswa_window=None, rswa_window_tensor=None, @@ -395,9 +395,9 @@ def fake_get_scheduler_metadata(**kwargs): assert torch.equal(builder.scheduler_metadata[:3], expected) -def test_update_draft_decode_metadata_skips_non_fa3_builders(monkeypatch): +def test_update_draft_decode_metadata_skips_without_scheduler_metadata(monkeypatch): builder = object.__new__(flash_attn_module.FlashAttentionMetadataBuilder) - builder.aot_schedule = False + builder.aot_schedule = True builder.use_full_cuda_graph = True builder.scheduler_metadata = torch.zeros(4, dtype=torch.int32) @@ -429,12 +429,12 @@ def fake_get_scheduler_metadata(**kwargs): num_prefill_reqs=0, num_decode_tokens=1, num_prefill_tokens=0, - scheduler_metadata=torch.tensor([5], dtype=torch.int32), + scheduler_metadata=None, prefix_scheduler_metadata=None, max_num_splits=1, causal=True, sliding_window=None, - mm_prefix_range_tensor=None, + mm_prefix_query_range_tensor=None, rswa_prefix_lens=None, rswa_window=None, rswa_window_tensor=None, @@ -443,7 +443,4 @@ def fake_get_scheduler_metadata(**kwargs): builder.update_draft_decode_metadata(metadata) assert not called - assert torch.equal( - metadata.scheduler_metadata, - torch.tensor([5], dtype=torch.int32), - ) + assert metadata.scheduler_metadata is None diff --git a/vllm/v1/attention/backends/flash_attn.py b/vllm/v1/attention/backends/flash_attn.py index 222b2c64b5d0..acdb1661dc95 100755 --- a/vllm/v1/attention/backends/flash_attn.py +++ b/vllm/v1/attention/backends/flash_attn.py @@ -803,7 +803,7 @@ def update_block_table( return new_metadata def update_draft_decode_metadata(self, metadata: FlashAttentionMetadata) -> None: - if not self.aot_schedule: + if metadata.scheduler_metadata is None: return num_reqs = metadata.num_decode_reqs or metadata.seq_lens.shape[0]