Skip to content

[FlyDSL] perf(pa_decode): add dynamic KV work planning and tune partition scheduling - #5546

Merged
fsx950223 merged 4 commits into
flydsl-pa-decode-tilefrom
zhiding512/flydsl-pa-dynamic-plan
Sep 16, 2026
Merged

fsx950223 merged 4 commits into
flydsl-pa-decode-tilefrom
zhiding512/flydsl-pa-dynamic-plan

Conversation

@zhiding512

@zhiding512 zhiding512 commented Sep 15, 2026 •

Copy link
Copy Markdown
Contributor

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

  • Add plan_pa_decode and pa_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.
  • Pack partial outputs by task and reduce only each request's actual partition count. Cover empty contexts, causal MTP tails, sparse physical pages, non-power-of-two partition limits, and poisoned scratch.
  • Keep planned execution explicit: uniform and short contexts can be slower. The torch.ops.aiter.pa_decode_flydsl interface remains static.
  • Allow the static split recommendation to use an optional host-known maximum context length, up to 256 partitions for small long-context batches. The CLI uses that hint; --max-partitions 8 retains the legacy clamp and --num-partitions remains an exact override.
  • Reduce live V-register overlap for fused MTP4 with page-128 caches and wide addresses. Page-16 and narrow-address paths retain their existing prefetch schedule.
  • Keep variable-context reference helpers and planner coverage in the existing PA unit test, document the Python API, and run the planner suite in the standard PA CI entry.

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.

Context lengths Static (us) Plan + attention + reduce (us) Speedup
B8: one 200003 + seven 257 97.28 45.19 2.15x
B200: one 200003 + 199 x 257 924.04 106.99 8.64x
B200: all 200000 3038.23 3245.32 0.94x
B8: all 257 15.40 20.13 0.76x

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.
  • Original command: all four page/layout combinations pass at atol=rtol=0.005.
  • Variable-context performance cases pass the causal FP32 reference and static/planned comparisons at atol=rtol=0.005.
  • Black, Ruff, git diff --check, and bash -n .github/scripts/aiter_test.sh pass.
python op_tests/test_flydsl_pa_decode.py -d bf16 -b 200 -q 4 \
  -s 16,1,128,200000 --block-size 16 128 --trans-v 0 1 --per-token 1

Usage and scratch layouts: docs/flydsl_pa_decode_plan.md.

@zufayu
zufayu requested a review from yadaish September 16, 2026 01:52
@fsx950223
fsx950223 merged commit cecd335 into flydsl-pa-decode-tile Sep 16, 2026
3 checks passed
@fsx950223
fsx950223 deleted the zhiding512/flydsl-pa-dynamic-plan branch September 16, 2026 07:34
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.

2 participants