[AMD] Use exact CU share for gfx950 segment-plan headroom - #39503
Conversation
b9ce3e0 to
c9ce1bd
Compare
The length-aware planner caps each sequence at twice its equal-share CU budget, but inherited the legacy power-of-two rounding. At batch discontinuities this leaves skewed long requests under-parallelized even though the planner's total workgroup budget prevents oversubscription. Derive the cap from the exact share while leaving the fixed split untouched, and pin the cap boundaries with CPU tests. Co-authored-by: Cursor <cursoragent@cursor.com>
c9ce1bd to
45ba783
Compare
The cap had been split into a private _seg_plan_max_segments so the CPU test could pass a CU count, leaving the public entry point an undocumented one-line wrapper and the test asserting against a private symbol. Take an optional num_cus on mtp_verify_attn_seg_max instead, matching how the other device-dependent launch-geometry heuristics here are written (_kv_splits_heuristic, _mla_split_budget) -- a CPX partition exposes 32 of an MI355X's 256 CUs, so an explicit count is useful beyond the test. The computed cap is unchanged. Move the test next to the other attention geometry contracts, have it drive the public helper, and drop the monkeypatch of module state. Co-authored-by: Cursor <cursoragent@cursor.com>
|
@zijiecode Can you help review this PR? |
|
@chuyeh Thank you for the improvement! Can we sweep more hicache points with 3600s complete perf? |
Sure. Happy to run this. Action items:
|
yichiche
left a comment
There was a problem hiding this comment.
This PR stops rounding the gfx950 MTP-verify attention planner’s per-sequence segment cap down to a power of two, so a long request in a skewed batch keeps its CU share instead of losing half its parallelism at batch-size discontinuities. LGTM.
@zijiecode Added the HiCache performance table to the PR. Let me know if you like to see any additional data points. |
Thank you, LGTM. |
sogalin
left a comment
There was a problem hiding this comment.
Dropping the pow2_floor only raises the per-sequence ceiling; the planner's total
workgroup budget is unchanged. LGTM.
Motivation
On gfx950 (MI355X), the assembly attention kernel from #37465 serves EAGLE/MTP verify by splitting each sequence's KV cache into segments, one workgroup each; the segment count sets how much of the GPU a sequence can use. #39172 made a length-aware planner the default for batch size > 1, but it still caps any one sequence at twice the legacy fixed-split count:
The
pow2_flooris inherited from the fixed split, which needs a power-of-two grid; the planner does not. The rounding wastes headroom at any batch size that is not a clean divisor. On a 256-CU MI355X with one KV head, going from 16 to 17 sequences halves the cap from 32 to 16 (the exact share of 15 rounds down to 8), so a long request can use only 16 of the planner's 256 workgroups while short batch mates leave the device idle.This PR uses the exact equal share instead:
Only the per-sequence ceiling rises; the planner's total workgroup budget is unchanged, so the device cannot be oversubscribed. Uniform batches already saturate the GPU and are unaffected; the headroom helps skewed batches where a few long requests own most of the KV traffic.
Related work applies the same CU-aware-cap idea to other split-KV verify paths: #35521 raises the flat split cap on the gfx942 Triton kernel, and #39316 ports split-KV verify to CUDA. This PR is the gfx950 assembly-planner counterpart.
Modifications
vattn_asm_gfx950/__init__.py:mtp_verify_attn_seg_maxuses the exact equal-share CU budget, not a power-of-two floor of it. An optionalnum_cus(default: current device) keeps the cap device-independent — a CPX partition exposes 32 of 256 CUs — and testable without a GPU.test/registered/unit/layers/attention/test_vattn_asm_segment_plan.py(new, CPU CI): pins the batch discontinuity, both GQA layouts, the 16/64 clamps, and monotonicity.4 lines of logic. The legacy fixed split, Triton fallback, planner/assembly kernels, kernarg ABI, and HIP-graph cache are untouched. Behavior changes only on gfx950 with the planner active (batch > 1).
Accuracy Tests
A wrong cap produces wrong output, not just a slowdown, so correctness is the real gate. The existing
test_vattn_segplan_mi35x.pycompares planned-split output against an fp32 reference and the legacy split, across both GQA ratios (16, 8), skewed/tiny/ragged lengths, batch 1/2/24/64, plus plan contents and the plan cache. The new cap exercises non-power-of-two strides (e.g. 20, 30) the old formula could never produce.MI355X, ROCm 7.2 —
7 passed, covering the new CPU contracts,test_vattn_segplan_mi35x.py, andtest_unified_attention_3d_mtp.py(asm vs aiter parity). Non-gfx950 hardware and the planner-inactive path do not read this cap and are unchanged.Speed Tests and Profiling
MI355X (gfx950), ROCm 7.2, head_dim 256, fp8 page-16 KV. Both arms are current
mainafter #39172; onlyseg_maxchanges. 50-100 launches, 5+ repeats, baseline bracketed around the candidate.Context sweep for one long request in bs17: neutral to 8k, then 1.21x/1.40x/1.62x at 16k/32k/70k; 1k within 2.2% on both GQA layouts. Trade-off: the fp32 partial buffers grow ~17 → 32 MiB at the bs17 q_len=4 16-head shape; the cap of 64 is unchanged and uniform kernel time stays flat.
End-to-end serving (InferenceX agentic replay)
Both tables use the MI355X InferenceX recipe replaying the public
cc-traces-weka-062126-256kdataset (seed 42), Qwen3.5-397B-A17B-MXFP4, TP2, concurrency 20,SGLANG_ASM_VERIFY_ATTN=1. Each tree runs once per GPU pair and the two runs are averaged, so per-pair speed differences cancel; both arms run concurrently on the same host. The cap was confirmed live in every container (16 base / 30 PR at batch 17). Comparing merge base2cb51f5d22with this PR.HBM only, 900 s/run (4 runs, 0 errors):
HiCache DRAM tier, 3600 s/run, balanced two-round crossover (4 runs, 12040 req, 0 errors):
mainRead these as "no regression", not a speedup. The deltas are smaller than run-to-run noise: within one arm, TPOT p90/p95 vary 6-9% across GPU pairs, and a third HiCache round flipped the tail signs (p90 +2.1%, p95 +5.5%). That is expected here — 55% of decode steps run at batch sizes where both formulas return an identical cap, and the rest only benefit when a sequence is long enough for the ceiling to bind. The kernel A/B above is the performance evidence; the E2E result shows this workload has no serving regression.
Reproduction note: HiCache on current
mainneedssgl_kernel.kvcacheio.get_device_accessible_ptr(#35233, 14 Sep); older images abort at startup on ROCm withKV_OFFLOADING=dram. These runs usedrocm/sgl-dev:v0.5.19-rocm720-mi35x-20260915.Checklist
CI States
Latest PR Test (Base): ✅ Run #35053158534
Latest PR Test (Extra): ❌ Run #35053158338
Latest PR Test (AMD ROCm 10): ❌ Run #35053158521