Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions docs/design/model_runner_v2.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
166 changes: 166 additions & 0 deletions tests/v1/worker/test_gpu_autoregressive_speculator.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@

from contextlib import nullcontext
from types import SimpleNamespace
from unittest.mock import Mock

import pytest
import torch
Expand All @@ -16,6 +17,9 @@
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
from vllm.v1.worker.gpu.spec_decode.autoregressive.speculator import (
Expand Down Expand Up @@ -278,3 +282,165 @@ 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(
(
"method_name",
"cg_mode",
"expected_eager_calls",
"expected_graph_replays",
),
[
("_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,
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)
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,
)

getattr(speculator, method_name)(
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_update_draft_decode_metadata_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_query_range_tensor=None,
rswa_prefix_lens=None,
rswa_window=None,
rswa_window_tensor=None,
)

builder.update_draft_decode_metadata(metadata)

assert torch.equal(metadata.scheduler_metadata, expected)
assert torch.equal(builder.scheduler_metadata[:3], expected)


def test_update_draft_decode_metadata_skips_without_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(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=None,
prefix_scheduler_metadata=None,
max_num_splits=1,
causal=True,
sliding_window=None,
mm_prefix_query_range_tensor=None,
rswa_prefix_lens=None,
rswa_window=None,
rswa_window_tensor=None,
)

builder.update_draft_decode_metadata(metadata)

assert not called
assert metadata.scheduler_metadata is None
4 changes: 4 additions & 0 deletions vllm/models/deepseek_v4/amd/rocm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
13 changes: 13 additions & 0 deletions vllm/v1/attention/backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -633,6 +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
# 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__(
Expand Down Expand Up @@ -757,6 +760,16 @@ def build_for_drafting(
fast_build=True,
)

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(
self,
common_prefix_len: int,
Expand Down
Loading
Loading