Skip to content

[Test] Enable CUDA coverage for shared PagedAttention contrib-op tests via cache aliasing harness - #31687

Merged
Hariharan Seshadri (hariharans29) merged 29 commits into
mainfrom
hari/wip-pagedattention-cuda-alias-tests
Aug 12, 2026
Merged

Hariharan Seshadri (hariharans29) merged 29 commits into
mainfrom
hari/wip-pagedattention-cuda-alias-tests

Conversation

@hariharans29

@hariharans29 Hariharan Seshadri (hariharans29) commented Aug 6, 2026 •

Copy link
Copy Markdown
Member

Description

Follow-up to #31611. This PR is intended to merge after #31611 because it builds on the shared contrib-op test relocation.

Please see #31611 (comment)

Problem

The shared OpTester-based suite uses separate output buffers, but CUDA PagedAttention requires cache output tensors to alias the corresponding cache input tensors.
As a result, CUDA path validation in the shared suite fails due to test harness semantics, not kernel correctness.

Goal

Enable CUDA execution for the shared contrib-op PagedAttention tests by adding an aliasing-capable test path that matches CUDA’s runtime contract.

Scope

  • Keep shared operator-level tests in contrib_ops.
  • Preserve existing WebGPU coverage (including non-aliased fallback behavior).
  • Add CUDA coverage using an IO-binding based harness where cache input/output share the same underlying buffer.
  • Avoid changing CUDA kernel functional behavior or aliasing requirements in this PR.

Out of Scope

  • Performance tuning or backend algorithm changes.
  • Broad refactors of existing PagedAttention test matrices.
  • Any schema or feature-expansion work unrelated to CUDA cache aliasing test enablement.

Validation Plan

  • Run targeted shared PagedAttention C++ tests on WebGPU (regression guard).
  • Run new CUDA aliasing-based tests and verify pass/fail behavior matches expectations.
  • Confirm CI passes for affected CUDA test legs.

Status

WIP branch created; implementation and targeted CUDA test harness wiring are in progress.

Motivation and Context

PagedAttention C++ op tests

Registers a NOT_IMPLEMENTED PagedAttention kernel for the WebGPU EP and lands the design doc describing the phased delivery plan. Follow-up PRs will implement the K/V writer, decode, and gather-then-flash prefill paths.
The helper is pure host code with no CUDA dependencies. Move it to contrib_ops/cpu/bert/ so it can be shared by other execution providers (CPU, WebGPU) without an EP-scope-violating include across contrib_ops/cuda/.
Replace the Phase 0 unconditional NOT_IMPLEMENTED with the full ComputeInternal control flow, minus the actual kernel launches:

- Fetch all 10 inputs and route them through the shared paged_attention_helper::CheckInputs, populating a PagedAttentionParameters.

- Populate the three non-helper fields (local_window_size, do_rotary, rotary_interleaved) from constructor state, matching the CUDA implementation.

- Enforce the do_rotary => cos_cache && sin_cache invariant with a specific error.

- Allocate output 0 with shape (token_count, hidden_size) and the two optional cache outputs with the paged shape (num_blocks, block_size, kv_num_heads, head_size).

- Enforce the schema-declared alias between input caches and output caches at compute time via a raw-pointer equality check (matches CUDA; no Alias/MayInplace on the KernelDef for now).

- Fast-path token_count == 0 to Status::OK.

- Branch the final NOT_IMPLEMENTED into distinct decode-vs-prefill messages that reference the design doc phase, so failures are informative.

Phase 1b (upcoming) will replace the two NOT_IMPLEMENTED tails with real WGSL kernel dispatch. See docs/design/webgpu_paged_attention.md §5.
… (Phase 1b.1)

Adds the first per-program CUDA-parity kernel for the WebGPU PagedAttention op: a plain (non-packed, non-rotary) scatter of new K/V tokens into the block-based paged cache.

* onnxruntime/contrib_ops/webgpu/bert/paged_attention_scatter_kv.wgsl.template: WGSL template. One invocation per (token, kv_head, dim); linear-scan cumulative_sequence_length to find seq_idx, then abs_slot = past_seqlens[seq] + local_tok, block_id = block_table[seq, abs_slot/block_size], slot = abs_slot%%block_size.

* onnxruntime/contrib_ops/webgpu/bert/paged_attention.h: adds ScatterKVToPagedCacheProgram with 8 Uint32 uniforms.

* onnxruntime/contrib_ops/webgpu/bert/paged_attention.cc: wires the scatter program from ComputeInternal, adds .MayInplace(3,1).MayInplace(4,2) hints on the KernelDef for the aliased GenAI fast path, and a copy-fallback for the non-aliased OpTester path (mirrors GroupQueryAttention). Output tensor is zero-filled until the attention path lands in Phase 1b.3/1b.4.

* onnxruntime/test/providers/webgpu/paged_attention_test.cc: 3 gtest cases covering single-token/no-past, multi-token with past, and multi-batch/multi-head with per-sequence past lengths and non-contiguous block_table.

Phase 1b.1 of docs/design/webgpu_paged_attention.md.
… (Phase 1b.2)

Adds a rotary embedding WGSL program used by the WebGPU PagedAttention op

to rotate Q and K in the non-packed layout before scattering K/V into the

paged cache. Mirrors paged_attention_impl.cu::RotaryEmbeddingTNH: same

interleaved-vs-split math, same position_id = past_seqlens[b] + s formula,

and dims >= rotary_dim are copied through unchanged.

ComputeInternal flow when do_rotary=1:

  1. Rotate query into output(0) (temporary layering until 1b.3 attention

     lands and overwrites output with real attention results).

  2. Rotate key into a GPU temp tensor.

  3. Scatter rotated key + untouched value into the paged cache via the

     existing ScatterKVToPagedCacheProgram from Phase 1b.1.

Value is not rotated. Packed-QKV + rotary path still returns NOT_IMPLEMENTED

(deferred to Phase 1b.2b).

Adds 3 gtests covering full-head non-interleaved, full-head interleaved,

and rotary_dim < head_size tail pass-through with multi-batch + GQA broadcast.

All 6 WebGpuPagedAttention.* tests pass.
Adds packed-QKV support to the WebGPU PagedAttention op. When `key`
and `value` are absent and the `query` input carries all three
projections concatenated per token (cols `[0, Q_hidden)` = Q,
`[Q_hidden, Q_hidden + KV_hidden)` = K, `[Q_hidden + KV_hidden,
Q_hidden + 2*KV_hidden)` = V), a new pre-pass kernel splits the
packed tensor into three standalone Q, K, V tensors. The rest of the
existing 1b.2 pipeline (optional non-interleaved / interleaved rotary
followed by paged-KV scatter) then runs unchanged against the split
tensors.

Design: split-then-reuse is intentionally conservative for the first
packed-QKV cut. It costs one extra full-tensor read/write in device
memory per Q/K/V column relative to a fused approach, but avoids
templating every downstream kernel on a packed-input layout and keeps
the CPU-side output-shape and cache-mutation reasoning identical to
the non-packed path. A fused rotary+scatter+packed variant can be
revisited in Phase 1c when we have baseline perf numbers.

Implementation:
- `PagedAttentionSplitPackedQKVProgram` (new): one WGSL kernel, one
  invocation per input element. Dispatch is
  `ceil(token_count * packed_hidden_size / WORKGROUP_SIZE)` groups.
  Uniforms carry `token_count`, `q_hidden_size`, `kv_hidden_size`,
  `packed_hidden_size`, `dispatch_size`.
- `paged_attention_split_packed_qkv.wgsl.template` (new): row-major
  linearization of `(token, packed_col)`, branching on the column
  range to route each element to the correct output tensor.
- `PagedAttention::ComputeInternal` (edited): when
  `parameters.is_packed_qkv` is true, allocate three transient GPU
  tensors of shapes `(token_count, hidden_size)`,
  `(token_count, kv_hidden_size)`, `(token_count, kv_hidden_size)`,
  run the split kernel, and rebind `query`/`key`/`value` locally to
  the split outputs before falling through to the existing rotary +
  scatter path.

Tests: extends the WebGPU PagedAttention test harness with a
`bool is_packed` field on both `ScatterCase` and `RotaryCase`, a
`PackQKV` helper (per-token concatenation of the reference float
buffers), and three new tests exercising the packed path:
`PackedQKV_NoRotary_MultiToken_SingleBatch`,
`PackedQKV_Rotary_NonInterleaved_SingleToken`,
`PackedQKV_Rotary_Interleaved_MultiBatch_GQA`. All 9
`WebGpuPagedAttention.*` tests pass.
…FA seqlens_q

Wires up the WebGPU PagedAttention kernel end-to-end for
continuous-batching / variable-Q-length workloads. Replaces the earlier
Phase 1a stub / Phase 1b.1-1b.2b sub-kernel scaffolding with the
production dispatch path:

    scatter K/V into paged cache
      -> gather paged K/V into padded BNSH scratch (RunGatherKV)
      -> unpack packed varlen Q into LEFT-aligned BSNH scratch
         (RunUnpackQuery)
      -> ApplyFlashAttention over padded scratch
      -> repack padded output back to (token_count, hidden_size)
         (RunRepackOutput)

## FlashAttention: optional seqlens_q input

The existing FA shader clamps
past_sequence_length = total_kv_b - max_seqlen_q to 0 on underflow.
That clamp is only correct for LEFT-aligned Q with past=0 (the GQA
"BatchedRightPaddedRotaryPrefill" scenario). For PagedAttention's
continuous-batching regime, past_b can be > 0 while q_len_b <
max_seqlen_q, and the clamp silently under-counts past_len_b, causing
real Q tokens to leak future KV positions through the causal mask
(observed as 85% mismatch in the s=16 packed=True test).

Introduces an optional per-batch new-Q-length input `seqlens_q` to
FA:

- `FlashAttentionProgram` / `FlashAttentionDecodeQKVProgram` gain a
  `use_seqlens_q_` template-conditional gate + `seqlens_q` shader
  input.
- When set, the shader computes
  past_sequence_length_b = total_kv_b - seqlens_q[b] = past_len_b
  which is always non-negative and correct for any (past, q_len)
  combination.
- Non-PA callers (GQA / MHA / Attention) pass nullptr, leave
  `use_seqlens_q_ = false`, and the shader takes the `#else` branch
  that is byte-identical to the pre-patch clamp path. Zero regression
  risk.
- `use_seqlens_q_` is included in the CacheHint for both programs to
  avoid pipeline-cache collision.

## PagedAttention: LEFT-aligned Q layout

`RunUnpackQuery` now places real tokens at padded slots [0, q_len_b)
with padding at [q_len_b, max_seqlen_q). `RunRepackOutput` mirrors
by reading from s = local_tok directly. This matches GQA's convention
and enables the correct per-batch past_len_b via seqlens_q above.

## Test coverage

- **32 / 32** WebGPU parity configs pass in
  `TestPagedAttentionWebGpu` (batch_size in {1,2}, sequence_length
  in {1,4,16}, MHA + GQA, packed on/off, block_size=256). The
  previously-failing test 25 (mixed q_len + past > 0) now passes.
- **5 / 5** C++ end-to-end tests
  (`WebGpuPagedAttention.EndToEnd_*`), including
  `EndToEnd_MixedPrefillDecode_MultiBatch_VariablePast`.
- **31 / 31** `GroupQueryAttention` WebGPU tests, including both
  `BatchedRightPaddedRotaryPrefill_WebGPU` and
  `BatchedRightPaddedRotaryPrefillFlashAttention_WebGPU`, unchanged
  since GQA doesn't pass seqlens_q.

## Cleanup: removed transitional Phase 1b.1 / 1b.2 / 1b.2b scaffolding

- Removed `_debug_mode` schema attribute + all three mode
  branches (unpack roundtrip, gather-slice verification, and
  legacy output=zeros/rotated_q).
- Removed `PagedAttentionGatherVerifyProgram` + its .wgsl.template
  + `RunGatherVerify`.
- Deleted 13 transitional gtests (`ScatterOnly_*`, `Rotary_*`,
  `PackedQKV_*`, `DebugMode_*`). The 5 `EndToEnd_*` tests cover the
  same functionality end-to-end; Python
  `TestPagedAttentionWebGpu` covers non-Linux platforms.

## Not in scope (deferred)

- `softcap != 0`: rejected with NOT_IMPLEMENTED.
- `local_window_size != -1`: rejected with NOT_IMPLEMENTED.
- `T = bfloat16`: only MLFloat16 registered.
- Graph capture (attention_metadata): documented as Phase 2 in
  `docs/design/webgpu_paged_attention.md` §4.4.
- Quantized KV cache (T_CACHE), MLA / LATENT, head_sink / QK-Norm:
  Phase 3 / 4 items from the design doc, tracked alongside CUDA
  parity work.

## Follow-up work (later PRs)

- Rewrite C++ Rotary_* and PackedQKV_* transitional tests to
  compare against an end-to-end reference so their coverage is
  restored on non-Linux CI.
- Add coverage-gap tests for `block_size != 256`, empty query
  (`token_count == 0`), and explicit non-default `scale`.
- Softcap + local_window_size in FlashAttentionProgram (also lifts
  GQA's `CanApplyFlashAttention` bailouts).
- Wire `TestPagedAttentionWebGpu` into a WebGPU CI leg. Today the
  Python parity suite runs on zero CI legs: the two WebGPU legs
  (linux_webgpu.yml, windows_webgpu.yml) are build-only, and
  nightly_webgpu.yml / macos-ci run `--test` but not
  `--enable_transformers_tool_test`. The C++
  `WebGpuPagedAttention.EndToEnd_*` gtests DO run on
  nightly_webgpu (Windows A10) and macos-ci (Metal), which is
  where CI protection sits today. A ~10-LOC follow-up to
  nightly_webgpu.yml can add a targeted pytest step for this file.
- paged_attention_test.cc: add missing #include <limits> (uses std::numeric_limits<float>::infinity()).

- paged_attention.cc: convert ORT_ENFORCE on the two optional cache outputs into ORT_RETURN_IF with a clearer error message (the scatter kernel needs both outputs, even though the schema marks them Optional).

- paged_attention.cc: move the input-to-output cache copy above the token_count==0 fast path so that the non-aliased path (OpTester) leaves initialized cache outputs even when there is no scatter work to do.
Copilot review noted that the doc's Phase 0 section was labeled '(this PR)' but this PR actually delivers Phase 1. Update Phase 0 label to '(early commits in this PR)' and move the '(this PR)' marker to Phase 1, which is the final state delivered.
…ention

# Conflicts:
#	onnxruntime/test/python/transformers/test_paged_attention.py
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>

@github-actions github-actions Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

You can commit the suggested changes from lintrunner.

Comment thread onnxruntime/test/contrib_ops/paged_attention_op_test.cc Outdated
Comment thread onnxruntime/test/contrib_ops/paged_attention_op_test.cc Outdated
Comment thread onnxruntime/test/contrib_ops/paged_attention_op_test.cc Outdated
Comment thread onnxruntime/test/contrib_ops/paged_attention_op_test.cc Outdated
Comment thread onnxruntime/test/contrib_ops/paged_attention_op_test.cc Outdated
Comment thread onnxruntime/test/contrib_ops/paged_attention_op_test.cc Outdated
Comment thread onnxruntime/test/contrib_ops/paged_attention_op_test.cc Outdated
Comment thread onnxruntime/test/contrib_ops/paged_attention_op_test.cc Outdated
Comment thread onnxruntime/test/contrib_ops/paged_attention_op_test.cc Outdated
Comment thread onnxruntime/test/contrib_ops/paged_attention_op_test.cc Outdated
Comment thread onnxruntime/test/contrib_ops/paged_attention_op_test.cc Fixed
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
Co-authored-by: github-actions[bot] <41898282+github-actions[bot]@users.noreply.github.com>
@hariharans29 Hariharan Seshadri (hariharans29) changed the title WIP: Enable CUDA coverage for shared PagedAttention contrib-op tests via cache aliasing harness Enable CUDA coverage for shared PagedAttention contrib-op tests via cache aliasing harness Aug 7, 2026
@hariharans29 Hariharan Seshadri (hariharans29) changed the title Enable CUDA coverage for shared PagedAttention contrib-op tests via cache aliasing harness [Test] Enable CUDA coverage for shared PagedAttention contrib-op tests via cache aliasing harness Aug 7, 2026

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Pull request overview

This PR extends PagedAttention testability across EPs by adding an IO-binding-based aliasing harness (to satisfy CUDA’s cache in-place contract) and by wiring in WebGPU PagedAttention coverage, including the supporting WebGPU kernel implementation and a small FlashAttention enhancement for variable per-batch Q lengths.

Changes:

  • Add an IO-binding contrib-op gtest harness that can bind cache inputs/outputs to the same underlying buffer (enabling CUDA execution).
  • Implement and register the WebGPU EP com.microsoft::PagedAttention v1 path (gather-then-flash fallback) plus required WGSL programs.
  • Extend WebGPU FlashAttention to optionally consume seqlens_q for correct causal masking with LEFT-aligned variable-q_len callers (PagedAttention).

Reviewed changes

Copilot reviewed 18 out of 19 changed files in this pull request and generated 2 comments.

Show a summary per file
File Description
onnxruntime/test/python/transformers/test_paged_attention.py EP-parameterize parity harness and add WebGPU parity matrix; add CUDA cache-aliasing IO-binding path.
onnxruntime/test/contrib_ops/paged_attention_op_test.cc New shared contrib-op gtests, including IO-binding alias/non-alias harness and WebGPU end-to-end correctness cases.
onnxruntime/contrib_ops/webgpu/webgpu_contrib_kernels.cc Register WebGPU PagedAttention kernel.
onnxruntime/contrib_ops/webgpu/bert/paged_attention.h Declare WebGPU PagedAttention kernel + supporting WebGPU programs.
onnxruntime/contrib_ops/webgpu/bert/paged_attention.cc Implement WebGPU PagedAttention (scatter, gather, unpack, FlashAttention, repack) and feature guards.
onnxruntime/contrib_ops/webgpu/bert/paged_attention_unpack_query.wgsl.template WGSL kernel to unpack packed varlen Q into padded BSNH scratch.
onnxruntime/contrib_ops/webgpu/bert/paged_attention_split_packed_qkv.wgsl.template WGSL kernel to split packed QKV into separate Q/K/V tensors.
onnxruntime/contrib_ops/webgpu/bert/paged_attention_scatter_kv.wgsl.template WGSL kernel to scatter K/V into the paged cache.
onnxruntime/contrib_ops/webgpu/bert/paged_attention_rotary.wgsl.template WGSL rotary embedding implementation for packed varlen Q or K.
onnxruntime/contrib_ops/webgpu/bert/paged_attention_repack_output.wgsl.template WGSL kernel to repack padded FA output back into packed varlen layout.
onnxruntime/contrib_ops/webgpu/bert/paged_attention_pack_metadata.wgsl.template WGSL kernel to pack metadata for a single device-to-host readback.
onnxruntime/contrib_ops/webgpu/bert/paged_attention_gather_kv.wgsl.template WGSL kernel to gather paged KV cache into dense padded BNSH scratch.
onnxruntime/contrib_ops/webgpu/bert/flash_attention.wgsl.template Add optional seqlens_q path for causal-mask loop bound computation.
onnxruntime/contrib_ops/webgpu/bert/flash_attention.h Plumb optional seqlens_q through FlashAttention program APIs.
onnxruntime/contrib_ops/webgpu/bert/flash_attention.cc Bind optional seqlens_q and include it in cache hints / shader parameters.
onnxruntime/contrib_ops/webgpu/bert/flash_attention_decode_qkv.wgsl.template Add optional seqlens_q path for decode QKV causal masking.
onnxruntime/contrib_ops/cuda/bert/paged_attention.cc Switch to the provider-neutral paged_attention_helper header.
onnxruntime/contrib_ops/cpu/bert/paged_attention_helper.h Add shared input-validation helpers for PagedAttention (used by CUDA/WebGPU).
docs/design/webgpu_paged_attention.md Add design/roadmap documentation for WebGPU PagedAttention.
Suppressed comments (2)

onnxruntime/test/python/transformers/test_paged_attention.py:45

  • The comment above _EP_TO_ORT_DEVICE says it maps EP -> (torch_device_when_cuda_available, ort_iobinding_device), but the dict values are just the ORT device string. This is confusing given torch device selection is implemented by Config.torch_device.
    onnxruntime/test/python/transformers/test_paged_attention.py:495
  • key_cache_np/value_cache_np are always materialized via .detach().cpu().numpy(), but they are unused in the CUDA+quantized path (which binds device pointers directly). Avoiding these unconditional host copies can significantly reduce test runtime for large caches.

Comment thread onnxruntime/contrib_ops/webgpu/bert/paged_attention.cc Outdated
Comment thread onnxruntime/test/contrib_ops/paged_attention_op_test.cc

@tianleiwu Tianlei Wu (tianleiwu) left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Three correctness issues need addressing before enabling WebGPU PagedAttention broadly: kernel registration currently claims unsupported cache dtypes, schema-legal nodes without cache outputs fail at runtime, and short rotary caches can be indexed out of bounds. The gather/FlashAttention design and variable-length coverage look solid otherwise.

Comment thread onnxruntime/contrib_ops/webgpu/bert/paged_attention.cc
Comment thread onnxruntime/contrib_ops/webgpu/bert/paged_attention.cc Outdated
Comment thread onnxruntime/contrib_ops/webgpu/bert/paged_attention.cc

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Pull request overview

Copilot reviewed 2 out of 2 changed files in this pull request and generated no new comments.

@tianleiwu Tianlei Wu (tianleiwu) left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

The WebGPU type constraints, optional-input guards, cache fallback behavior, and rotary-cache validation look sound. I found one blocking coverage gap: the new CUDA path does not run the shared reference-backed cases, and its single smoke case can pass without verifying that the cache update occurred. Details inline.

Comment thread onnxruntime/test/contrib_ops/paged_attention_op_test.cc
Addresses PR #31687 review: the CUDA-aliased smoke test was checking only pointer identity and a non-zero output element. Since the value_cache is pre-filled with 0.02 and past_seqlen=4, a scatter regression could still leave output[0] non-zero, silently passing. Copy both bound cache outputs back and EXPECT_NEAR the past_seqlen slot to the scattered K/V values (0.03/0.04). Covers both aliased and non-aliased paths uniformly by reading via io_binding->GetOutputs().
@hariharans29
Hariharan Seshadri (hariharans29) merged commit 10008a6 into main Aug 12, 2026
86 of 87 checks passed
@hariharans29
Hariharan Seshadri (hariharans29) deleted the hari/wip-pagedattention-cuda-alias-tests branch August 12, 2026 00:51
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants