Repository navigation
[CUDA] Make adaptive FlashDecode splits CUDA-graph safe - #32239
Justin Chu (justinchuby) wants to merge 7 commits into
Conversation
## 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>
There was a problem hiding this comment.
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.
Tianlei Wu (tianleiwu)
left a comment
There was a problem hiding this comment.
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.
…uda-attention-capture-safe-split-kv
Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com>
Tianlei Wu (tianleiwu)
left a comment
There was a problem hiding this comment.
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.
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>
Tianlei Wu (tianleiwu)
left a comment
There was a problem hiding this comment.
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>
|
Justin Chu (@justinchuby), please resolve conflicts. |
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
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