Skip to content

[CUDA] Make adaptive FlashDecode splits CUDA-graph safe - #32239

Open
Justin Chu (justinchuby) wants to merge 7 commits into
microsoft:mainfrom
justinchuby:justinchu/automated/cuda-attention-capture-safe-split-kv
Open

Justin Chu (justinchuby) wants to merge 7 commits into
microsoft:mainfrom
justinchuby:justinchu/automated/cuda-attention-capture-safe-split-kv

Conversation

@justinchuby

Copy link
Copy Markdown
Contributor

Summary

Plan contiguous GQA FlashDecode split-KV launches from fixed KV-cache capacity during CUDA graph capture and replay, while retaining live-sequence-length planning for ordinary eager execution.

Why

CUDA graphs freeze launch geometry and workspace addresses at capture time. Planning NumSplits from the current live sequence length can become stale as the cache grows, preventing adaptive split-KV behavior from remaining valid across replay. Graph-enabled warmup now reserves capacity-sized workspace, capture uses the fixed-capacity plan, and eager decode avoids redundant capacity heuristic work.

Behavior

  • Capture/replay uses fixed cache capacity for stable NumSplits and workspace sizing
  • Graph warmup reserves replay-sized workspace before capture
  • Eager execution continues to tune from the live sequence length
  • Active memset size remains limited to the launch plan
  • Debug output reports the resolved NumSplits

Validation

Focused host tests cover head sizes 64, 128, and 256; local-window and sequence-tail behavior; non-decode inputs; capture planning; and ordinary eager routing. For an SM108 configuration with live length 129 and capacity 4097, capture selected 17/17/22 splits for head sizes 64/128/256, while eager retained live-length plans. Independent review found one redundant eager heuristic computation, which is fixed in this commit.

No CUDA kernel timing is claimed because this Windows host did not have nvcc. The change preserves eager routing and targets graph planning correctness and replay-stable adaptive split-KV behavior.

Based on the mechanisms validated in justinchuby/onnx-genai#1340 and #1350.

Description

Motivation and Context

## Summary

Plan contiguous GQA FlashDecode split-KV launches from fixed KV-cache capacity during CUDA graph capture and replay, while retaining live-sequence-length planning for ordinary eager execution.

## Why

CUDA graphs freeze launch geometry and workspace addresses at capture time. Planning NumSplits from the current live sequence length can become stale as the cache grows, preventing adaptive split-KV behavior from remaining valid across replay. Graph-enabled warmup now reserves capacity-sized workspace, capture uses the fixed-capacity plan, and eager decode avoids redundant capacity heuristic work.

## Behavior

- Capture/replay uses fixed cache capacity for stable NumSplits and workspace sizing
- Graph warmup reserves replay-sized workspace before capture
- Eager execution continues to tune from the live sequence length
- Active memset size remains limited to the launch plan
- Debug output reports the resolved NumSplits

## Validation

Focused host tests cover head sizes 64, 128, and 256; local-window and sequence-tail behavior; non-decode inputs; capture planning; and ordinary eager routing. For an SM108 configuration with live length 129 and capacity 4097, capture selected 17/17/22 splits for head sizes 64/128/256, while eager retained live-length plans. Independent review found one redundant eager heuristic computation, which is fixed in this commit.

No CUDA kernel timing is claimed because this Windows host did not have nvcc. The change preserves eager routing and targets graph planning correctness and replay-stable adaptive split-KV behavior.

Based on the mechanisms validated in justinchuby/onnx-genai#1340 and microsoft#1350.

Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>
Copilot AI balanced review requested due to automatic review settings August 24, 2026 18:19

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

Makes adaptive GQA FlashDecode split planning stable across CUDA graph capture and replay.

Changes:

  • Uses KV-cache capacity during capture while preserving live-length eager planning.
  • Reserves capture-sized workspace and limits active memset sizes.
  • Adds host-side split-planning tests and debug reporting.

Reviewed changes

Copilot reviewed 4 out of 4 changed files in this pull request and generated 1 comment.

File Description
onnxruntime/contrib_ops/cuda/bert/group_query_attention.cc Implements capture-aware split and workspace planning.
onnxruntime/contrib_ops/cuda/bert/group_query_attention.h Tracks CUDA graph enablement.
onnxruntime/contrib_ops/cuda/bert/group_query_attention_impl.h Declares the split-plan helper and result type.
onnxruntime/test/providers/cuda/test_cases/attention_split_heuristic_test.cc Tests host-side eager and capture plans.

💡 Configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

@justinchuby

Copy link
Copy Markdown
Contributor Author

Tianlei Wu (@tianleiwu)

@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 capture/replay split planning looks coherent, but the current head does not compile in the Linux CUDA and TensorRT checks because the new test exposes a missing header dependency. Please fix the build blocker before merging. I did not duplicate the existing end-to-end CUDA graph test thread, which remains open.

Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>

@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 capture-aware split planning looks coherent: warm-up reserves capacity-sized workspace, capture uses capacity-derived launch geometry, and replay bypasses ComputeInternal. One blocker remains in the new end-to-end test: its host cache data and reference use the wrong layout for the bound BNSH tensor, so the test does not validly exercise the change.

Comment thread onnxruntime/test/contrib_ops/group_query_attention_op_test.cc Outdated
Preserve both CUDA cache-aliasing coverage and the FlashDecode graph replay regression test.

Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>
Use head-major cache offsets for initialization, appended-value mirroring, and the output reference to match the bound BNSH tensor layout.

Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>
Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>

@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.

Re-reviewed the current head after the follow-up fixes. The prior header dependency and BNSH reference-indexing blockers are addressed, and the end-to-end test now exercises CUDA graph capture/replay with both KV cache pairs aliased and growing live sequence lengths. The capture-aware planner keeps active launch geometry separate from the maximum reserved workspace, while replay bypasses host ComputeInternal. I found no remaining actionable issues.

Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>
@tianleiwu

Copy link
Copy Markdown
Contributor

Justin Chu (@justinchuby), please resolve conflicts.

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

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants