Repository navigation
[Perf][Model] Use block-sparse attention in Qwen2.5-Omni Token2Wav DiT - #7975
linyueqian merged 2 commits into
Conversation
|
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. |
|
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:
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. |
9e2c85a to
2a8be5f
Compare
Omni ReviewBot: supersededThe CI failure noted on |
linyueqian
left a comment
There was a problem hiding this comment.
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) |
There was a problem hiding this comment.
[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) |
There was a problem hiding this comment.
[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>
2a8be5f to
e5cef35
Compare
|
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 Review fixes
Online serving numbers
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).
(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.) |
|
Gentle ping @linyueqian, the new head e5cef35 still has no Buildkite run (the |
linyueqian
left a comment
There was a problem hiding this comment.
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.
|
The failing CI does not seem to be related to this PR. |
Omni ReviewBot routing recordAssigned Strict on zcode (GLM-5.3-Flash) under experiment |
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 allnkeys. 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
block_sparse_attention()inqwen2_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/DiTDecoderLayerpassblock_sizeand the look-ahead/look-backward counts instead of a dense mask;_create_block_diffis removed.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).
pytest tests/model_executor/models/qwen2_5_omni/ -q. The newtest_token2wav_block_attention.pycomparesblock_sparse_attentionto 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.Qwen/Qwen2.5-Omni-3B, 3 requests per run, main then this PR back to back on the same GPU:e2e_stage_2_wall_time_msfrom--log-stats.use_mixed_modalitiesanduse_audio_in_video.Test Result
98 passed(9 new).Token2Wav time per request (all 3 requests shown; the first includes JIT warmup):
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.
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).
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, andtools/pre_commit/{check_test_marks,check_spdx_header,check_forbidden_imports,check_torch_cuda}.pyall 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.