Skip to content

[Spec Decode] Let the speculator build its own draft prefill attention metadata - #53660

Draft
LucasWilkinson wants to merge 1 commit into
mainfrom
lwilkinson/drafter-builds-step0-attn
Draft

LucasWilkinson wants to merge 1 commit into
mainfrom
lwilkinson/drafter-builds-step0-attn

Conversation

@LucasWilkinson

@LucasWilkinson LucasWilkinson commented Aug 25, 2026 •

Copy link
Copy Markdown
Contributor

Stack: 1. #53660 (this PR) -> 2. #53661 (PCP+MTP, based on this branch)

Purpose

The autoregressive speculator's prefill pass reuses the target model's attention metadata dict and slot mappings. That reuse is valid because of the identical padded batch layout and KV-cache slots, and it is paired with a capture-time coupling: draft prefill CUDA-graph capture builds its dummy metadata through the target runner's builders and buffers so capture and runtime stay consistent.

This cross-model contract breaks whenever a feature transforms the target batch between the target forward and the drafter. The immediate motivator is prefill context parallelism (#53427): under PCP the target batch is rank-sharded while the replicated drafter must see the global batch, so the target's metadata describes the wrong layout.

This PR adds the capability for the drafter to build its own prefill metadata, gated so the default path is unchanged:

  • AutoRegressiveSpeculator.prepare_attn() builds the draft prefill attention metadata and slot mappings from the input batch through the drafter's own attention groups, mirroring how the decode steps already build theirs (naming mirrors the model runner's / PCPManager's prepare_attn). It runs before FULL-graph replay as well, since builder.build() refreshes the persistent state the captured graph reads.
  • A new reuse_target_attn_metadata flag (default True) selects between reusing the target's metadata (existing behavior, zero added work) and building via prepare_attn(). Batch-transforming features clear the flag; [PCP][Spec Decode] Consolidate MTP support and add DSpark groundwork #53661 does so for the replicated-PCP drafter.
  • Prefill CUDA-graph capture follows the same flag, so capture and runtime always build through the same builders and buffers in either mode.
  • In prepare_attn(), slot mappings are computed from input_batch.positions rather than the drafter's positions buffer: prepare_prefill_inputs only copies accepted tokens, so the drafter's buffer is stale at rejected slots, which would emit KV writes into live cache slots.

With the flag at its default, this PR is behavior- and cost-neutral: no extra kernels, no extra metadata builds.

Why this is not duplicating existing work

#53427 adds equivalent metadata construction wired through the PCP manager, two model-runner call sites, and a widened _build_draft_attn_metadata signature. This PR was written from scratch as the outcome of maintainer review of #53427: it places the same construction where it architecturally belongs (the speculator), behind an explicit gate, so #53427's PCP support reduces to a small config/partitioning diff (see #53661, which credits the original author).

Test Plan

  • pytest tests/v1/worker/test_gpu_autoregressive_speculator.py tests/v1/spec_decode/test_eagle_draft_attn_metadata.py tests/v1/worker/test_gpu_extract_hidden_states_speculator.py
  • pre-commit run (ruff, mypy) on changed files
  • Default (reuse) path: unchanged from main; covered by existing spec-decode suites
  • prepare_attn() path: exercised under PCP in [PCP][Spec Decode] Consolidate MTP support and add DSpark groundwork #53661's GPU validation (GLM-4.7-Flash, MTP K=3, PCP4 — see that PR's results)

Test Result

  • Unit tests: 26 passed across the four suites above; lint/mypy: pass
  • Reuse path (default, B300, greedy, 8 prompts x 64 tokens): Qwen3-8B + RedHatAI/Qwen3-8B-speculator.eagle3 reproduces the pre-change run byte-for-byte (acceptance 0.3537, 738 drafted, identical outputs); GLM-4.7-Flash + MTP K=3 (MLA) acceptance 0.353 — statistically identical to the build-path run (0.350)
  • prepare_attn() path: validated under PCP in [PCP][Spec Decode] Consolidate MTP support and add DSpark groundwork #53661, including a GSM8K accuracy/acceptance comparison vs TP4 (see that PR's results)

AI assistance: this PR was implemented with Claude Code under my direction and review; I reviewed every changed line.

🤖 Generated with Claude Code

@mergify mergify Bot added speculative-decoding dflash mrv2 Model Runner V2 specific labels Aug 25, 2026
@LucasWilkinson
LucasWilkinson force-pushed the lwilkinson/drafter-builds-step0-attn branch from 7c2a155 to 65dbf86 Compare August 25, 2026 04:25
@LucasWilkinson LucasWilkinson changed the title [Spec Decode] Build draft step-0 attention metadata in the speculator [Spec Decode] Build draft prefill attention metadata in the speculator Aug 25, 2026
…n metadata

The autoregressive speculator's prefill pass reuses the target model's
attention metadata and slot mappings, relying on the identical padded
batch layout and on the runner building draft-layer metadata entries
through the same builders its cudagraph capture uses. This cross-model
contract breaks whenever a feature transforms the target batch between
the target forward and the drafter (e.g. prefill context parallelism,
where the target batch is rank-sharded while the drafter must see the
global batch).

Add prepare_attn(), which builds the draft prefill metadata through the
drafter's own attention groups from the input batch, mirroring how the
decode steps already build theirs, gated on a new
reuse_target_attn_metadata flag. The flag defaults to True (reuse, no
behavior or cost change); batch-transforming features clear it. Prefill
cudagraph capture follows the same flag so capture and runtime always go
through the same builders and buffers.

Slot mappings in prepare_attn are computed from input_batch.positions
rather than the drafter's positions buffer: prepare_prefill_inputs only
copies accepted tokens, so the drafter's buffer is stale at rejected
slots, which would emit KV writes into live cache slots.

Co-authored-by: Claude Fable 5 <noreply@anthropic.com>
Signed-off-by: Lucas Wilkinson <lwilkins@redhat.com>
@LucasWilkinson LucasWilkinson changed the title [Spec Decode] Build draft prefill attention metadata in the speculator [Spec Decode] Let the speculator build its own draft prefill attention metadata Aug 25, 2026
@LucasWilkinson
LucasWilkinson force-pushed the lwilkinson/drafter-builds-step0-attn branch from 65dbf86 to e921e36 Compare August 25, 2026 04:36
@tomasruizt

tomasruizt commented Sep 16, 2026 •

Copy link
Copy Markdown
Contributor

This functionality seems to be implemented for SD method=draft_model. There are two PRs for it: #43091, #55990. The latter references your design. Do you need any specific API contract to combine it with PCP down the line? @LucasWilkinson

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

Status: Backlog

Development

Successfully merging this pull request may close these issues.

2 participants