Skip to content
Draft
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
84 changes: 71 additions & 13 deletions vllm/v1/worker/gpu/spec_decode/autoregressive/speculator.py
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,45 @@ def advance_draft_positions(self) -> bool:
"""
return True

def prepare_attn(
self,
input_batch: InputBatch,
batch_desc: BatchExecutionDescriptor,
) -> tuple[dict[str, Any] | None, dict[str, torch.Tensor]]:
"""Build attention metadata and slot mappings for the draft prefill pass.

The draft prefill keeps the target batch's padded query layout, so the
metadata is built from the input batch through the drafter's own
attention groups rather than reusing the target model's metadata.
"""
num_reqs = input_batch.num_reqs
num_reqs_padded = batch_desc.num_reqs or num_reqs
self.block_tables.gather_block_tables(
input_batch.idx_mapping,
num_reqs_padded=num_reqs_padded,
)
# Slot mappings must use input_batch.positions: the drafter's own
# positions buffer is stale at rejected token slots, which would
# produce KV writes into live cache slots.
slot_mappings_tensor = self.block_tables.compute_slot_mappings(
input_batch.idx_mapping,
input_batch.query_start_loc,
input_batch.positions,
batch_desc.num_tokens,
)
slot_mappings = build_slot_mappings_by_layer(
slot_mappings_tensor, self.kv_cache_config
)
attn_metadata = self._build_draft_attn_metadata(
num_reqs=num_reqs,
num_reqs_padded=num_reqs_padded,
num_tokens_padded=batch_desc.num_tokens,
seq_lens_cpu_upper_bound=input_batch.seq_lens_cpu_upper_bound,
step=0,
query_start_loc_np=input_batch.query_start_loc_np,
)
return attn_metadata, slot_mappings

def set_attn(
self,
model_state: ModelState,
Expand Down Expand Up @@ -156,20 +195,26 @@ def capture(self) -> None:
# For FULL graphs, the entire routine is recorded as one graph.
# For PIECEWISE, only the model's compiled regions are captured
# and the rest (compute_logits, gumbel_sample) runs eagerly.
# Draft prefill reuses the target model's attention metadata at
# runtime, so capture builds its dummy metadata through the target
# model runner's builders and buffers.
# Capture must build its dummy metadata through the same builders and
# buffers the runtime metadata comes from: the target model runner's
# when reusing the target's metadata, the drafter's own otherwise.
assert self.prefill_cudagraph_manager is not None
if self.prefill_cudagraph_manager.use_breakable_cg:
self.prefill_cudagraph_manager.init_breakable_cg_runner(self.model)

if self.reuse_target_attn_metadata:
prefill_input_buffers = self.target_input_buffers
prefill_attn_groups = self.target_attn_groups
else:
prefill_input_buffers = self.input_buffers
prefill_attn_groups = self.attn_groups
self.on_prefill_begin(self.max_num_reqs)
self.prefill_cudagraph_manager.capture(
self._prefill,
self.model_state,
self.target_input_buffers,
prefill_input_buffers,
self.block_tables,
self.target_attn_groups,
prefill_attn_groups,
self.kv_cache_config,
progress_bar_desc="Capturing prefill CUDA graphs",
)
Expand Down Expand Up @@ -237,9 +282,12 @@ def propose(
# NOTE(woosuk): To avoid CPU-GPU synchronization without CPU knowing the
# number of rejected tokens, we maintain the size of input_ids and
# hidden_states the same as the target model's. This means, we pad each
# request's query length to include any rejected positions. By doing so,
# we can also reuse the attention metadata (e.g., query_start_loc,
# seq_lens) of the target model.
# request's query length to include any rejected positions. By doing
# so, we can also reuse the attention metadata (e.g., query_start_loc,
# seq_lens) of the target model — unless the target batch was
# transformed (reuse_target_attn_metadata is False), in which case
# prepare_attn rebuilds it over the same padded layout through the
# drafter's own attention groups.
if aux_hidden_states:
assert self.method == "eagle3"
hidden_states = self.model.combine_hidden_states(
Expand Down Expand Up @@ -289,6 +337,19 @@ def propose(
need_eager=is_profile,
)

# When the drafter cannot reuse the target's metadata, build it here —
# even for FULL replay, since builder.build() refreshes the persistent
# state the captured graph reads, mirroring the decode steps.
prefill_attn_metadata: dict[str, Any] | None = attn_metadata
prefill_slot_mappings: dict[str, torch.Tensor] | None = slot_mappings
if not self.reuse_target_attn_metadata:
if dummy_run and skip_attn_for_dummy_run:
prefill_attn_metadata, prefill_slot_mappings = None, None
else:
prefill_attn_metadata, prefill_slot_mappings = self.prepare_attn(
input_batch, prefill_batch_desc
)

self._prepare_eplb_forward(num_tokens)

self.on_prefill_begin(num_reqs)
Expand All @@ -297,14 +358,11 @@ def propose(
assert self.prefill_cudagraph_manager is not None
self.prefill_cudagraph_manager.run_fullgraph(prefill_batch_desc)
else:
# The target model's attention metadata and slot mappings
# can directly be used for draft prefill, because of the
# identical batch shape and KV cache layout.
self._prefill(
num_reqs,
prefill_batch_desc.num_tokens,
attn_metadata,
slot_mappings,
prefill_attn_metadata,
prefill_slot_mappings,
num_tokens_across_dp=num_tokens_across_dp,
cudagraph_runtime_mode=prefill_batch_desc.cg_mode,
mm_inputs=mm_inputs,
Expand Down
9 changes: 9 additions & 0 deletions vllm/v1/worker/gpu/spec_decode/speculator.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,6 +74,15 @@ def __init__(self, vllm_config: VllmConfig, device: torch.device):
self.vllm_config = vllm_config
self.device = device

# Reuse the target model's attention metadata and slot mappings for
# the draft prefill pass (identical padded batch layout and KV cache
# slots, so rebuilding them is redundant work). Features that
# transform the target batch between the target forward and the
# drafter (e.g. a PCP-sharded target with a replicated drafter) must
# clear this so the drafter builds its own metadata from the input
# batch; cudagraph capture follows the same choice of builders.
self.reuse_target_attn_metadata = True

assert vllm_config.speculative_config is not None
self.speculative_config = vllm_config.speculative_config
self.method = self.speculative_config.method
Expand Down
Loading