Skip to content

[AMD] Use exact CU share for gfx950 segment-plan headroom - #39503

Merged
HaiShaw merged 3 commits into
sgl-project:mainfrom
chuyeh:amd/vattn-asm-segment-fill
Sep 22, 2026
Merged

HaiShaw merged 3 commits into
sgl-project:mainfrom
chuyeh:amd/vattn-asm-segment-fill

Conversation

@chuyeh

@chuyeh chuyeh commented Sep 15, 2026 •

Copy link
Copy Markdown
Contributor

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:

seg_max = clamp(2 * pow2_floor(num_cus / (batch_size * num_kv_heads)), 16, 64)

The pow2_floor is 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:

seg_max = clamp(2 * floor(num_cus / (batch_size * num_kv_heads)), 16, 64)

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_max uses the exact equal-share CU budget, not a power-of-two floor of it. An optional num_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.py compares 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, and test_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 main after #39172; only seg_max changes. 50-100 launches, 5+ repeats, baseline bracketed around the candidate.

shape cap active WGs current main exact share speedup
hkv1, uniform bs17 x 70k 16 → 30 255 → 255 109.7 us 109.8 us 1.00x
hkv1, agentic skew bs17 16 → 30 49 → 83 238.0 us 134.0 us 1.78x
hkv1, agentic skew bs24 16 → 20 56 → 65 238.3 us 193.0 us 1.23x
hkv2, uniform bs9 x 70k 16 → 28 126 → 126 111.9 us 112.2 us 1.00x
hkv2, agentic skew bs9 16 → 28 41 → 69 223.1 us 139.1 us 1.60x
hkv1, one 70k request in bs17 16 → 30 32 → 46 75.6 us 46.7 us 1.62x

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-256k dataset (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 base 2cb51f5d22 with this PR.

HBM only, 900 s/run (4 runs, 0 errors):

TP2 c20, HBM Merge base This PR Delta
TPOT p50 10.42 ms 10.26 ms -1.6%
TPOT p90 17.86 ms 17.57 ms -1.6%
TPOT p95 24.74 ms 23.98 ms -3.1%
TTFT p50 1.063 s 1.012 s -4.8%
Output 705.5 tok/s 707.0 tok/s +0.2%

HiCache DRAM tier, 3600 s/run, balanced two-round crossover (4 runs, 12040 req, 0 errors):

TP2 c20, HiCache 1.5 Current main This PR Delta
TPOT p50 9.63 ms 9.58 ms -0.52%
TPOT p90 15.22 ms 14.90 ms -2.10%
TPOT p95 20.00 ms 19.57 ms -2.13%
TTFT p50 1.110 s 1.065 s -4.02%
Output 754.2 tok/s 762.2 tok/s +1.07%

Read 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 main needs sgl_kernel.kvcacheio.get_device_accessible_ptr (#35233, 14 Sep); older images abort at startup on ROCm with KV_OFFLOADING=dram. These runs used rocm/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

@chuyeh
chuyeh force-pushed the amd/vattn-asm-segment-fill branch from b9ce3e0 to c9ce1bd Compare September 15, 2026 03:24
@chuyeh chuyeh changed the title [AMD] Size the vattn_asm split-KV segments to the launch geometry [AMD] Stop rounding the vattn_asm split-KV segment count down to a power of two Sep 15, 2026
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>
@chuyeh
chuyeh force-pushed the amd/vattn-asm-segment-fill branch from c9ce1bd to 45ba783 Compare September 15, 2026 07:57
@chuyeh chuyeh changed the title [AMD] Stop rounding the vattn_asm split-KV segment count down to a power of two [AMD] Use exact CU share for gfx950 segment-plan headroom Sep 15, 2026
@chuyeh
chuyeh marked this pull request as ready for review September 15, 2026 12:23
@yichiche yichiche added the run-ci CI: run the baseline test suite on this PR label Sep 16, 2026
yichiche and others added 2 commits September 16, 2026 10:04
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>
@yichiche

Copy link
Copy Markdown
Collaborator

@zijiecode Can you help review this PR?

@zijiecode

Copy link
Copy Markdown
Collaborator

@chuyeh Thank you for the improvement! Can we sweep more hicache points with 3600s complete perf?

@chuyeh

chuyeh commented Sep 16, 2026 •

Copy link
Copy Markdown
Contributor Author

@chuyeh Thank you for the improvement! Can we sweep more hicache points with 3600s complete perf?

Sure. Happy to run this.

Action items:

  • 3600 s runs, HiCache DRAM tier (KV_OFFLOADING=dram KV_OFFLOAD_BACKEND=hicache), same crossover method as above
  • TP2 concurrency 20 as a null control (caps are identical for bs >= 29, so it should show zero difference)

@yichiche yichiche left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

@chuyeh

chuyeh commented Sep 17, 2026

Copy link
Copy Markdown
Contributor Author

@chuyeh Thank you for the improvement! Can we sweep more hicache points with 3600s complete perf?

Sure. Happy to run this.

Action items:

  • 3600 s runs, HiCache DRAM tier (KV_OFFLOADING=dram KV_OFFLOAD_BACKEND=hicache), same crossover method as above
  • TP2 concurrency 20 as a null control (caps are identical for bs >= 29, so it should show zero difference)

@zijiecode Added the HiCache performance table to the PR. Let me know if you like to see any additional data points.

@zijiecode

Copy link
Copy Markdown
Collaborator

@chuyeh Thank you for the improvement! Can we sweep more hicache points with 3600s complete perf?

Sure. Happy to run this.

Action items:

  • 3600 s runs, HiCache DRAM tier (KV_OFFLOADING=dram KV_OFFLOAD_BACKEND=hicache), same crossover method as above
  • TP2 concurrency 20 as a null control (caps are identical for bs >= 29, so it should show zero difference)

@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 sogalin left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Dropping the pow2_floor only raises the per-sequence ceiling; the planner's total
workgroup budget is unchanged. LGTM.

@HaiShaw
HaiShaw merged commit c2f14bf into sgl-project:main Sep 22, 2026
302 of 348 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

jit-kernel run-ci CI: run the baseline test suite on this PR

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants