[FlyDSL] perf(pa_decode): add dynamic KV work planning and tune partition scheduling - #5546
Merged
fsx950223 merged 4 commits intoSep 16, 2026
Merged
Conversation
Merged
1 task
4 of 5 tasks
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Requests with very different KV lengths receive the same static partition count, leaving a few long requests on the critical path. Add an opt-in GPU planner that apportions partitions by actual context length and launches packed KV work. For one 200003-token request mixed with short 257-token requests, the measured plan + attention + reduction latency improves by 2.15x at batch 8 and 8.64x at batch 200.
Targets
flydsl-pa-decode-tile(#4332), which provides the underlying PA implementation.Changes
plan_pa_decodeandpa_decode(..., work_plan=plan). The Triton planner reads GPU-resident lengths, assigns contiguous 256-token tile ranges, and writes reusable task/reduction metadata. Attention and reduction remain FlyDSL. Refreshing preallocated plans needs no CPU readback and supports graph replay.torch.ops.aiter.pa_decode_flydslinterface remains static.--max-partitions 8retains the legacy clamp and--num-partitionsremains an exact override.Performance
Measured on gfx950 with 256 CUs. BF16 query/output, MTP4, Hq16/Hkv1/D128, FP8 K/V with per-token scales, page128, transposed V. Comparisons below use identical sparse-page inputs and the repository's 101-iteration timer with allocation rotation. Dynamic latency includes a plan refresh on every call plus attention and reduction. Static and planned columns refer to the two paths in this PR.
These results support using dynamic plans for uneven KV work, not enabling them universally.
For the original B200/C200000 command, the four static cases report 3.43-4.32 TB/s effective bandwidth and 426-535 TFLOP/s, approximately 124 FLOP/byte. These are logical-traffic estimates that count referenced KV tokens once and exclude repeated loads and scratch; they are not hardware HBM-counter measurements or a measured percentage of peak. The rocprofv3 diagnostic traces for mixed contexts recorded no scratch spills. No full hardware-counter roofline measurement is claimed.
Validation
python -m pytest -q op_tests/test_flydsl_pa_decode.py: 281 passed, including the base branch's MTP2/MTP3 cases, static high-partition/wide-address cases, dynamic planning, graph replay, and other reducer dimensions.atol=rtol=0.005.atol=rtol=0.005.git diff --check, andbash -n .github/scripts/aiter_test.shpass.Usage and scratch layouts:
docs/flydsl_pa_decode_plan.md.