Skip to content

[Perf][Model] Use block-sparse attention in Qwen2.5-Omni Token2Wav DiT - #7975

Merged
linyueqian merged 2 commits into
vllm-project:mainfrom
LiRunGuo:perf/qwen25-omni-token2wav-block-attention
Oct 8, 2026
Merged

linyueqian merged 2 commits into
vllm-project:mainfrom
LiRunGuo:perf/qwen25-omni-token2wav-block-attention

Conversation

@LiRunGuo

Copy link
Copy Markdown

Purpose

After the thinker and talker, Qwen2.5-Omni's Token2Wav stage (flow-matching DiT + BigVGAN) is the slowest part of a request: about 10–20 s on an H200 for 30–40 s of audio. Almost all of that is the DiT (22 layers, fp32, 36 forward passes per request: 10-step RK4 with CFG run as batch 2); BigVGAN takes ~0.3 s.

The DiT only attends inside a window of 24-frame mel blocks. 19 of the 22 layers attend to their own block only, and the other 3 add one neighbour block (look_backward_layers / look_ahead_layers). But every layer built a dense (batch, heads, n, n) mask from an int64 block-difference tensor and ran dense SDPA over all n keys. With torch.profiler on one 4000-frame sample, about 35% of the CUDA time is the mask ops (ge/le/bitwise_and/where) and about 30% is the dense attention kernel, where each query keeps only 24–48 of the 4000 keys.

Change

  • Add block_sparse_attention() in qwen2_5_omni_token2wav.py: reshape q/k/v into blocks, attend to the own block plus the shifted neighbour blocks, and mask only the padded positions of a partial last block. The math is the same as the dense masked SDPA; the work is linear in sequence length instead of quadratic.
  • DiTAttention / DiTDecoderLayer pass block_size and the look-ahead/look-backward counts instead of a dense mask; _create_block_diff is removed.
  • No config, weight, or API change. Token2Wav stays fp32 (bf16 and TF32 do not help here; the cost is the mask traffic, not the GEMMs).

Test Plan

vLLM Version: 0.29.0 (torch 2.13.0+cu130, CUDA 13.0, driver 580.159.03)

vLLM-Omni Commit: f8a00b1 (main) vs this PR

Hardware: NVIDIA H200 141 GB. For the A/B, all three stages run on one GPU so both variants see identical placement; the GPU had no other processes during the runs (checked before and after each run).

  1. Unit tests (CPU): pytest tests/model_executor/models/qwen2_5_omni/ -q. The new test_token2wav_block_attention.py compares block_sparse_attention to dense masked SDPA for seq_len 10 / 96 / 101 (shorter than a block, multiple of the block size, partial last block) and for plain, look-backward, and look-ahead layers.
  2. End-to-end A/B, Qwen/Qwen2.5-Omni-3B, 3 requests per run, main then this PR back to back on the same GPU:
    cd examples/offline_inference/qwen2_5_omni
    python end2end.py --model Qwen/Qwen2.5-Omni-3B --query-type <use_video|use_mixed_modalities> \
      --num-prompts 3 --log-stats --stage-overrides '{"0": {"enforce_eager": false, "devices": "0", "gpu_memory_utilization": 0.4}, "1": {"enforce_eager": false, "devices": "0", "gpu_memory_utilization": 0.3}, "2": {"devices": "0", "gpu_memory_utilization": 0.15}}'
    The overrides only put all stages on one GPU and turn on cudagraph for the AR stages (see [Performance]: Qwen2.5-Omni thinker/talker run eager on CUDA; enabling cudagraph gives ~4-5x faster decode on H200 #7960) so the timeline is less noisy. Token2Wav itself runs the same way in both variants. Token2Wav time is e2e_stage_2_wall_time_ms from --log-stats.
  3. Output check: thinker text and talker token count compared per request, WAV-to-WAV correlation between main and this PR, and Whisper-large-v3 transcription of every WAV compared to the thinker text.
  4. Smoke run with the bundled deploy YAML (no overrides, 2x H200) for use_mixed_modalities and use_audio_in_video.

Test Result

  1. 98 passed (9 new).

  2. Token2Wav time per request (all 3 requests shown; the first includes JIT warmup):

    query type audio per request main this PR
    use_video 26.8 s 10.06 / 10.34 / 8.61 s 4.20 / 4.74 / 2.48 s
    use_mixed_modalities 40.3 s 18.60 / 20.67 / 20.71 s 4.87 / 4.79 / 3.23 s

    Total wall time for the 3 requests: 31.2 s → 20.6 s (use_video), 54.3 s → 27.8 s (use_mixed_modalities). With all stages on one GPU, requests 1 and 2 overlap with the next request's thinker/talker, so the third request is the cleanest number (3.5x and 6.4x). Longer audio gains more because the dense path is quadratic in length.

    Standalone DiT sample, 4000 mel frames, fp32, same GPU: 10.8 s → 2.87 s; peak extra memory 2.74 GiB → 0.78 GiB.

  3. Same thinker text and talker token count for every request; same audio duration; WAV correlation ≥ 0.99982 between main and this PR; Whisper similarity 1.000 for both. The standalone mel difference is at most 6e-6 (fp32 rounding from a different reduction order).

  4. Both pass (rc=0, audio produced). Token2Wav took 3.28 s for the mixed-modalities request with the bundled YAML.

Lint on the changed files: ruff check, ruff format --check, typos, and tools/pre_commit/{check_test_marks,check_spdx_header,check_forbidden_imports,check_torch_cuda}.py all pass.

Not measured: concurrency above 1 (the Token2Wav stage uses max_num_seqs: 1, so this is a per-request latency change), the 7B model, and hardware other than H200. Qwen3-TTS's 25 Hz tokenizer (modeling_qwen3_tts_tokenizer_v1.py) builds the same dense block mask; I did not touch or measure it here.

AI assistance: I used Claude Code to profile the stage, draft the change and tests, run the A/B, and draft this description. I reviewed the changes and checked the numbers against the logs.

@vllm-omni-review-bot

Copy link
Copy Markdown

This PR appears to belong to: docs/design/module/model_integration.md, docs/design/module/ar_runtime.md.

Module owners: @tzhouam @fake0fan @Gaohan123

Routing: @tzhouam via module of the changed files, CODEOWNERS; @fake0fan via module of the changed files; @Gaohan123 via module of the changed files

@LiRunGuo, please review your own changes and leave a short self-review comment describing what you checked. PRs without author self-review may not be assigned a reviewer.

Please take a look when you have a chance. If you would like an automated review, mention @vllm-omni-review-bot in a comment.

@LiRunGuo

Copy link
Copy Markdown
Author

Self-review:

I found this when profiling Qwen2.5-Omni on H200. After the talker, Token2Wav was the slowest stage, and most of its time was not the real attention compute. Every DiT layer built a full n x n mask from the block-diff tensor and then ran dense SDPA, but each query only really looks at 24 or 48 keys (its own block, plus one neighbour block on 3 layers).

What I checked:

  • The new block_sparse_attention gives the same result as the old dense masked SDPA. I tested short sequence (smaller than one block), exact multiple of block size, and a partial last block, for normal / look-backward / look-ahead layers. 98 tests pass in tests/model_executor/models/qwen2_5_omni/ (9 new).
  • A/B on one H200, main vs this PR back to back on the same GPU, 3 requests each. I checked there were no other processes on the GPU during the runs, because our server is shared. Token2Wav goes from ~10 s to ~2.5-4.7 s for use_video, and ~20 s to ~3.2-4.9 s for use_mixed_modalities.
  • Output is the same: same thinker text and talker tokens, same audio length, WAV correlation >= 0.9998 vs main, and Whisper transcript matches the text (1.000) for both.
  • After rebasing on [Core] Optimize Qwen2.5-Omni Token2Wav buffer loading #7061 (it also touched this file, only buffer loading), I ran the unit tests again and a smoke run with the default deploy YAML. Both fine.

Not tested: concurrency > 1 (this stage is max_num_seqs 1 anyway), 7B, and GPUs other than H200. Qwen3-TTS 25Hz tokenizer has the same dense-mask code, I did not change it here, maybe a follow-up.

@linyueqian linyueqian added the ready label to trigger buildkite CI label Sep 22, 2026
@hsliuustc0106 hsliuustc0106 added the enhancement New feature or request label Sep 23, 2026
@LiRunGuo
LiRunGuo force-pushed the perf/qwen25-omni-token2wav-block-attention branch from 9e2c85a to 2a8be5f Compare September 23, 2026 00:28
@linyueqian linyueqian added ready label to trigger buildkite CI and removed ready label to trigger buildkite CI labels Sep 23, 2026
@vllm-omni-review-bot

vllm-omni-review-bot commented Sep 23, 2026 •

Copy link
Copy Markdown

Omni ReviewBot: superseded

The CI failure noted on 2a8be5fa89d7 refers to an earlier head; the pull request now points at e5cef357b5b0.

@linyueqian linyueqian left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

The dense block mask is replaced by attention over a local window of key blocks (own block plus look_backward_block before and look_ahead_block after), built by shifting the blocked key and value tensors per offset, masking the shifted-out and padded positions, and running one scaled_dot_product_attention over the concatenated window. Static read of the indexing, the padding of the last block, the validity mask and the output reshape found no correctness regression, and the nine new unit cases cover the block boundaries, padding and the equivalence with the dense path. Two optional performance notes are inline: the bool(valid.all()) mask decision forces a device sync per call even though the parameters already decide it, and the offset-zero iteration allocates and fills full shifted copies that are discarded in the common own-block-only case. Static read at 2a8be5fa against merge-base f8a00b14, both files; fork head, no PR code executed; pre-commit and DCO green, the general lane's only failures are the inherited test_wan_vae_fastpath_install.py cases from the #7056 break, fixed on main by #8047. This approval covers correctness only; as a perf change it should show an end-to-end gain (Token2Wav latency or RTF on a real serving run) before it merges.

values = torch.cat(values, dim=3)
valid = torch.cat(valids, dim=1)
# Every query block keeps at least its own block, so no row is fully masked.
mask = None if bool(valid.all()) else valid.view(1, 1, num_blocks, 1, -1)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

[suggestion] bool(valid.all()) is a device-to-host sync on every attention call. Whether any position is masked is already known from the parameters: the mask can only contain False when pad > 0 or when a neighbour block is shifted past either end (look_backward_block > 0 or look_ahead_block > 0). Deciding mask = None from those Python ints keeps the hot path free of the sync.

for offset in range(-look_backward_block, look_ahead_block + 1):
# Key block i + offset for each query block i; blocks past either end
# are zero-filled and masked out.
shifted_key = torch.zeros_like(key)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

[suggestion] For offset zero this allocates zeros_like copies of key, value and the validity mask and fills them, only to discard them when len(keys) == 1. Handling the own-block-only case before the loop (reuse key, value and key_valid directly) removes three allocations and copies from the most common configuration.

The Token2Wav DiT attends within a sliding window of 24-frame mel blocks
(own block, plus one neighbour block on the look-ahead/look-backward
layers), but every layer built a dense (batch, heads, n, n) mask from an
int64 block-difference tensor and ran dense SDPA over all n keys. The
mask construction alone was about a third of the stage's CUDA time.

Compute attention per block instead: reshape q/k/v into blocks and attend
to the own block plus the shifted neighbour blocks, masking only padded
positions of a partial last block. The result matches the dense masked
SDPA up to fp32 rounding, and work is linear in sequence length instead
of quadratic.

Signed-off-by: RunguoLi <li19107254665@gmail.com>
…in Token2Wav attention

Decide whether block_sparse_attention needs a mask from the padding and
look-ahead/look-backward counts instead of reading it back from the GPU
with valid.all(), and attend to the blocked keys directly for own-block
layers instead of building shifted copies that were then discarded. The
own-block offset of neighbour layers also reuses the blocked keys.

Signed-off-by: RunguoLi <li19107254665@gmail.com>
@LiRunGuo
LiRunGuo force-pushed the perf/qwen25-omni-token2wav-block-attention branch from 2a8be5f to e5cef35 Compare September 29, 2026 20:06
@LiRunGuo

Copy link
Copy Markdown
Author

Thanks for the review @linyueqian! I took both suggestions in e5cef35 (separate commit so it's easy to see), and rebased on main so the Wan VAE failure (#8047) should be gone now.

One small request: the push didn't start the Buildkite lanes on the new head (the ready label was already on, #7609), and I can't change labels. Could you remove and re-add ready when you have a moment? The GitHub checks (pre-commit, build 3.11/3.12, DCO, docs) are already green on e5cef35.

Review fixes

  • The mask decision now comes from the Python ints (pad, look_backward_block, look_ahead_block), no valid.all() any more. I checked with torch.cuda.set_sync_debug_mode("error"): the old version raised on every call, the new one does not.
  • Own-block layers (19 of 22) use the blocked key/value directly, and the offset-0 block of the neighbour layers also reuses them instead of a zero-filled copy.
  • Output is bit-identical to the previous commit (max diff 0.0 on GPU, B=2, H=16, n=4000/4001, d=64). Per call on H200: own-block 0.40 -> 0.30 ms, neighbour layers 0.70 -> 0.60 ms. tests/model_executor/models/qwen2_5_omni/ 98 passed.

Online serving numbers

vllm serve Qwen/Qwen2.5-Omni-7B --omni with the bundled deploy YAML (no overrides, so thinker/talker still eager), vLLM 0.30.0, 2x H200 per server. Benchmark is vllm bench serve --omni with the same workload as the Qwen3-Omni perf CI (random-mm, 1 audio item, input/output 100, --ignore-eos, --num-warmups 2):

vllm bench serve --omni --model Qwen/Qwen2.5-Omni-7B --backend openai-chat-omni --endpoint /v1/chat/completions \
  --dataset-name random-mm --random-input-len 100 --random-output-len 100 --random-range-ratio 0.0 --ignore-eos \
  --random-mm-base-items-per-request 1 --random-mm-num-mm-items-range-ratio 0.5 \
  --random-mm-limit-mm-per-prompt '{"audio":1}' --random-mm-bucket-config '{"(0, 60, 3)":1.0}' \
  --percentile-metrics ttft,tpot,itl,e2el,audio_rtf,audio_ttfp,audio_duration \
  --num-warmups 2 --max-concurrency <1|4> --request-rate inf --num-prompts <8|16>

Our GPU server is shared and the CPU load from other jobs changed a lot during my first try (load average ~100 -> ~170), which made that run meaningless. So I ran main and this PR at the same time on two GPU pairs, with every benchmark hitting both servers together, then swapped the GPU pairs and did it again. Both variants produced exactly the same work (same thinker/talker token counts and audio frames per request).

concurrency 1 (8 req) concurrency 4 (16 req)
mean audio RTF, main -> PR 1.72 -> 1.60 (0.930), 1.73 -> 1.56 (0.901) 15.8 -> 15.0 (0.949), 15.5 -> 13.1 (0.849)
mean e2el, main -> PR 26.1 -> 25.3 s, 28.2 -> 25.1 s 41.5 -> 38.7 s, 43.8 -> 39.8 s
request throughput, PR/main 1.030, 1.124 1.017, 1.066

(two numbers per cell = round 1, round 2.) Averaged over the rounds: mean audio RTF -8.5% at concurrency 1 and -10% at concurrency 4, mean e2el -7% / -8%, throughput +8% / +4%. 0 failed requests.

The end-to-end gain is smaller than the Token2Wav stage speedup in the offline runs in the description, because with the default eager config the talker decode takes most of each request. With cudagraph on for the AR stages (#7960), Token2Wav should be a bigger share of the request, but I did not measure that combination online here. Also the RTF > 1 here is because Qwen2.5-Omni returns the audio only at the end (no async chunk), so RTF includes the whole pipeline.

(Claude Code helped me run these benchmarks and draft this reply; I checked the numbers against the result JSONs.)

@LiRunGuo

LiRunGuo commented Oct 5, 2026

Copy link
Copy Markdown
Author

Gentle ping @linyueqian, the new head e5cef35 still has no Buildkite run (the ready label was already on when I pushed, so nothing fired, #7609). Could you or another maintainer (cc @Gaohan123) remove and re-add ready so the GPU lanes run? The review fixes and the online serving numbers are in my comment above. Thanks!

@linyueqian linyueqian added ready label to trigger buildkite CI and removed ready label to trigger buildkite CI labels Oct 5, 2026

@linyueqian linyueqian left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Both earlier suggestions are addressed in e5cef357b: the mask decision uses the Python block and padding parameters, own-block layers reuse the blocked key/value tensors, and offset zero also reuses them in the neighbour path. The padding and shifted-neighbour masks retain their prior behavior, and I found no new correctness issues in this delta.

The online 7B serving measurements in your follow-up address the earlier request for end-to-end evidence. Because these are author-reported results on vLLM 0.30 and shared H200 hardware, they do not establish a vLLM 0.31 performance result. Main CUDA for this head: build 16705.

Static review of the delta from 2a8be5fa, the unchanged nine unit cases, and the affected surface on main 0b0d2d700; no PR code executed.

@BeatSeat

BeatSeat commented Oct 6, 2026

Copy link
Copy Markdown
Contributor

The failing CI does not seem to be related to this PR.
LGTM.

@vllm-omni-review-bot

Copy link
Copy Markdown
Omni ReviewBot routing record

Assigned Strict on zcode (GLM-5.3-Flash) under experiment fleet-strict-cursor-grok46-zcode-glm53flash-5050-c5-z10-20261002.

@linyueqian
linyueqian merged commit 59196c6 into vllm-project:main Oct 8, 2026
7 of 9 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

enhancement New feature or request ready label to trigger buildkite CI

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants