[WebGPU] Optimized PagedAttention implementation (2/n) (#31727)
## [WebGPU] PagedAttention: direct paged decode, fused paged prefill,
Unpack/Repack skip (Phase 2 partial)
This PR is the Phase 2 follow-up to #31611. It replaces the "always
gather + always Unpack/Repack" v1 fallback with two paged-aware
FlashAttention programs that read the paged KV cache directly, and a
fast path that lets FA consume the packed varlen Q buffer without
materializing padded BSNH scratch. Net effect: **~2× faster decode,
~1.15× faster uniform prefill, ~1.3× faster varlen prefill** on the
shape matrix below, with no regressions.
The Phase 1 gather-then-flash fallback shipped in #31611 remains intact
and still runs on adapters / configs where the paged-aware shaders can't
safely dispatch (see `Correctness invariants` below).
### What's shipped
1. **Direct paged split-reduce decode.** `FlashAttentionPagedDecodeQKV`
+ `FlashAttentionPagedDecodeVxReduce` index `key_cache` / `value_cache`
directly through `block_table`. Selected when `max_seqlen_q < 32` —
mirrors the dense-FA split-reduce threshold. Eliminates the dense K/V
scratch and its gather bandwidth for every decode step.
2. **Fused paged prefill.** `FlashAttentionPagedPrefillProgram` is a
straight port of the dense-FA prefill shader's shared-memory path with
page-table-aware K/V tile loads
(`bert/flash_attention_paged_prefill.wgsl.template`). Supports fp16,
BSNH Q, packed varlen Q (`q_varlen` template variant), and
variable-Q-length causal masking via `seqlen_k` + `seqlens_q`. No
attention_bias / head_sink / TurboQuant.
3. **Unpack/Repack skip fast paths.** When direct paged attention runs,
we can hand FA a rank-4 view over the raw packed Q buffer instead of
allocating padded BSNH scratch:
- **Uniform mode** (`B * max_seqlen_q == token_count`): view is `[B,
max_seqlen_q, N, H]`. Covers decode, `B==1` prefill, and equal-length
batched prefill (the common continuous-batching case).
- **Varlen mode**: view is `[token_count, 1, N, H]` plus
`cumulative_seqlens_q`; only the fused paged-prefill shader can index it
(`q_varlen`).
Skipping Unpack+Repack removes 2 dispatches (~300–500 µs of CPU dispatch
cost per Run on D3D12) plus a `B * max_seqlen_q * hidden * 2 B` scratch
allocation (tens of MB at long prefill).
### Dispatch-count reduction
| Route (no rotary, non-packed) | #31611 (merged) | This PR (shm-path
adapters) |
|---|---|---|
| **Decode** (`max_seqlen_q < 32`) | Scatter + Gather + UnpackQ +
DecodeQKV + DecodeVxReduce + Repack = **6** | Scatter + PagedDecodeQKV +
PagedDecodeVxReduce = **3** |
| **Prefill** (`max_seqlen_q ≥ 32`) | Scatter + Gather + UnpackQ +
FlashAttention + Repack = **5** | Scatter + FlashAttentionPagedPrefill =
**2** |
Decode's FA is 2 kernels (split-K: QKV + VxReduce); prefill's FA is 1
kernel (`FlashAttentionProgram`). On configs where
`ShouldRunFusedPagedPrefill` rejects (fp32, `head_size > 256`,
`block_size < max_k_step`), the prefill row falls back to the 5-dispatch
#31611 cascade; decode's direct paged split-reduce path has no such
gate. Neither route is adapter-gated — the paged shaders use no subgroup
intrinsics and run on every WebGPU adapter that meets the fp16 /
shm-budget / alignment predicates.
### Correctness invariants
**Prefill selection** consults one shared predicate:
```cpp
bool ShouldRunFusedPagedPrefill(context, is_fp16, max_seqlen_q,
head_size, block_size);
```
It rejects (→ gather-then-flash fallback) when any of:
- `!is_fp16` — only fp16 variant is compiled today.
- `max_seqlen_q < 32` — decode uses the split-reduce programs instead.
- `head_size` exceeds the workgroup shared-memory budget (fp16:
`head_size > 256`).
- `block_size < max_k_step` — the fused shader assumes one K/V tile
lives in one paged block (one `block_table` lookup per tile).
`paged_attention_helper` only enforces `block_size >= 16` power-of-two;
e.g. `block_size=16` with fp16 `head_size<=128` (`max_k_step=32`) would
splice into a physically-adjacent block that isn't the next entry in the
table.
The fused paged-prefill shader uses only workgroup shared memory (no
subgroup intrinsics), so there is no adapter-class gate — subgroup
adapters (Qualcomm / AMD / Intel with subgroups) take the paged shm
kernel directly instead of falling back to gather + dense-FA-subgroup.
Because the same predicate gates the "skip `RunGatherKV`", "skip
`q_padded` scratch", and "select fused shader" decisions, the three
cannot drift.
**Decode selection** (`max_seqlen_q < 32`) is a pure shape check — no
adapter, dtype, or block-size gate. The direct paged split-reduce
kernels (`FlashAttentionPagedDecodeQKV` +
`FlashAttentionPagedDecodeVxReduce`) are the sole decode path when the
kernel dispatches at all (fp16 is enforced at kernel registration, so no
fp32 fallback is possible). Unlike fused prefill, the decode kernels do
one `block_table` lookup per K/V slot rather than per tile, so they have
no `block_size` alignment requirement.
**WGSL correctness gotcha handled** in the fused prefill shader.
`cumulative_seqlens_q` is `array<i32>` but row indices are `u32`. Both
`loadq` and `writeo` explicitly cast (`u32(cumulative_seqlens_q[b]) +
q_idx`); without the cast, tint surfaces the type-resolution failure as
an opaque `absl::…raw_hash_map<>::at` at runtime.
### Performance
Machine: dev-box discrete WebGPU adapter (D3D12), 24-core host. Google
Benchmark harness at
`onnxruntime/test/onnx/microbenchmark/paged_attention.cc`,
`--benchmark_min_time=0.3s`, wall-clock timing via `UseManualTime()`.
Earlier revisions of this PR included an
`ORT_WEBGPU_PAGED_ATTENTION_USE_FUSED` env-var kill switch used for A/B
measurement against the #31611 cascade. That toggle has been removed
(the direct/fused paths are selected internally by shape and config; the
numbers below are the reason). The A/B was performed by temporarily
broadening the toggle locally to also force
`use_direct_paged_decode=false` and `skip_unpack_repack=false`, so
fused=0 exercised the exact gather-then-flash cascade shipped in #31611.
All numbers below are with that broadened toggle; the broadening was
reverted before final push.
Column meanings: **nH** = num query heads, **nKV** = num KV heads, **H**
= head dim. Shape families:
- MHA_H64 (nH=16, nKV=16, H=64), MHA_H128 (nH=16, nKV=16, H=128)
- GQA_Qwen (nH=14, nKV=2, H=128), GQA_Llama (nH=32, nKV=4, H=128)
#### Decode (16 shapes)
| Shape (B/nH/nKV/H/past) | this PR (µs) | #31611 (µs) | Speedup |
|---|---:|---:|---:|
| 1/16/16/128/2048 | 669 | 3612 | **5.40×** |
| 2/16/16/128/512 | 627 | 3248 | **5.18×** |
| 1/16/16/64/2048 | 861 | 3835 | **4.45×** |
| 2/16/16/64/2048 | 907 | 3617 | **3.99×** |
| 2/16/16/64/512 | 603 | 1369 | 2.27× |
| 2/16/16/128/2048 | 2129 | 4816 | 2.26× |
| 1/16/16/128/512 | 607 | 1150 | 1.89× |
| 2/32/4/128/2048 | 1590 | 3010 | 1.89× |
| 1/14/2/128/2048 | 979 | 1601 | 1.64× |
| 2/32/4/128/512 | 656 | 990 | 1.51× |
| 1/16/16/64/512 | 576 | 843 | 1.46× |
| 2/14/2/128/512 | 694 | 984 | 1.42× |
| 1/14/2/128/512 | 600 | 757 | 1.26× |
| 2/14/2/128/2048 | 970 | 1194 | 1.23× |
| 1/32/4/128/512 | 643 | 751 | 1.17× |
| 1/32/4/128/2048 | 1434 | 1437 | 1.00× |
Range **1.00×–5.40×**, geomean ~2.0×. Biggest wins on long-past MHA
(H=128, past=2048) where gather bandwidth dominated. The one 1.00× row
is a small-K/V-cache GQA case where gather cost was already low.
#### Uniform prefill (24 shapes)
Range **1.01×–1.25×**, geomean ~1.13×. Highlights (all wins):
| Shape (B/nH/nKV/H/T) | this PR (µs) | #31611 (µs) | Speedup |
|---|---:|---:|---:|
| 1/32/4/128/128 | 1225 | 1530 | **1.25×** |
| 2/14/2/128/128 | 1111 | 1385 | **1.25×** |
| 2/16/16/64/128 | 766 | 948 | 1.24× |
| 2/32/4/128/512 | 9723 | 12054 | 1.24× |
| 1/16/16/128/128 | 856 | 1051 | 1.23× |
| 2/16/16/128/1024| 17296 | 21029 | 1.22× |
| 1/14/2/128/128 | 718 | 877 | 1.22× |
| 2/32/4/128/128 | 1640 | 1966 | 1.20× |
| 1/16/16/64/128 | 750 | 891 | 1.19× |
| 2/16/16/128/128 | 1294 | 1491 | 1.15× |
*(14 more rows 1.01×–1.15×; full log in tree.)*
Short-T shapes gain most from Unpack/Repack skip; long-T shapes are
dominated by FA compute time.
#### Varlen prefill (12 shapes, halving q_lens = `{max_T, max_T/2,
max_T/4, …}`)
Range **1.14×–1.73×**, geomean ~1.29×.
| Shape (B/nH/nKV/H/maxT) | q_lens | this PR (µs) | #31611 (µs) |
Speedup |
|---|---|---:|---:|---:|
| 4/16/16/128/512 | `{512,256,128,64}` | 5987 | 10375 | **1.73×** |
| 4/14/2/128/512 | `{512,256,128,64}` | 5312 | 7815 | **1.47×** |
| 4/32/4/128/512 | `{512,256,128,64}` | 11165 | 15427 | **1.38×** |
| 4/16/16/128/1024 | `{1024,512,256,128}` | 20127 | 27661 | **1.37×** |
| 2/32/4/128/512 | `{512,256}` | 8345 | 10732 | 1.29× |
| 2/16/16/128/1024 | `{1024,512}` | 15121 | 19084 | 1.26× |
| 4/14/2/128/1024 | `{1024,512,256,128}` | 17685 | 21562 | 1.22× |
| 4/32/4/128/1024 | `{1024,512,256,128}` | 39571 | 47468 | 1.20× |
| 2/14/2/128/1024 | `{1024,512}` | 13141 | 15430 | 1.17× |
| 2/16/16/128/512 | `{512,256}` | 4528 | 5235 | 1.16× |
| 2/32/4/128/1024 | `{1024,512}` | 28871 | 33383 | 1.16× |
| 2/14/2/128/512 | `{512,256}` | 4067 | 4645 | 1.14× |
Wins grow with batch size — bigger B means more of the padded-BSNH
round-trip gets eliminated (B=4/maxT=512 packs only 46.9% of `B·maxT`
tokens; the padded scratch #31611 allocates is >2× bigger than the
actual data).
### Tests
- `onnxruntime/test/contrib_ops/paged_attention_op_test.cc`
`PagedAttention.EndToEnd_*` — 12/12 non-CUDA tests pass. Covers MHA,
GQA, single/multi-batch, variable past lengths, empty tokens, packed
QKV, rotary, mixed prefill+decode, cache aliasing via IO-binding. New:
`EndToEnd_Prefill_MultiBatch_Varlen_Fused` (B=2, token_count=48 with
q_lens (32,16), head_size=128, MHA) — regression test for the fused
varlen prefill path.
- Micro-benchmark harness
`onnxruntime/test/onnx/microbenchmark/paged_attention.cc` — 52
registered shapes (16 decode + 24 uniform prefill + 12 varlen prefill).
### Related
- Phase 1 (v1 fallback): #31611 (merged)
- Schema extensions: #29912 (merged)
- Design doc: `docs/design/webgpu_paged_attention.md` (updated in this
PR)
---------
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
Co-authored-by: Tianlei Wu <tlwu@microsoft.com>