Skip to content

[AMD][gfx950] wo_a route above 64 rows, quantise once, and MXFP8 tile fixes - #10

Open
tianxiaojiang4 wants to merge 1 commit into
kevin-mii:dsv41-amd-mainfrom
tianxiaojiang4:amd-gfx950-decode-projections
Open

tianxiaojiang4 wants to merge 1 commit into
kevin-mii:dsv41-amd-mainfrom
tianxiaojiang4:amd-gfx950-decode-projections

Conversation

@tianxiaojiang4

@tianxiaojiang4 tianxiaojiang4 commented Sep 22, 2026 •

Copy link
Copy Markdown

Four measured changes to the gfx950 decode projections, stacked on dsv41-amd-main (the branch behind sgl-project#39857).
Every one is gated against the unmodified branch and validated on a live server. Opened as a draft against your
branch rather than main, since all four files are yours and not upstream yet — happy to fold them into sgl-project#39857
instead if you would rather carry them.

The changes

1. wo_a above the split-K ceiling. Past _SPLIT_K_MAX_M the 16×32 grid tile re-streams each group's whole
weight once per 16 rows — about 800 MB of traffic for a 16 MiB weight. A strided-batched bf16 GEMM written
straight into the [T, G, R] layout the consumer already wants, with the same fp8-grid rounding applied once
afterwards, takes it from 125.3 → 33.1 µs at 768 rows. SGLANG_OPT_WO_A_LARGE_M_ROUTE=triton restores the
old tile. Below 64 rows nothing changes.

2. Quantise once. On that route the result is rounded onto the fp8 grid in its own pass, and wo_b — which
sees no input scale — then quantises the same 2048-wide activation again. Emitting Mxfp8Activation instead
removes one of the two passes and is bitwise identical; wo_b's own quantise kernel disappears because the
operand now arrives with a scale. Measured on the pair, four alternating processes: 31.6 / 37.4 / 43.3 / 56.2
→ 28.0 / 32.9 / 38.2 / 50.3 µs at M = 96 / 192 / 384 / 768, a flat −12 %.
SGLANG_OPT_WO_A_LARGE_M_EMIT=grid restores the current emission.

3. Skip the tile masks when the tile divides the problem. They are loop-invariant, and the disassembly shows
26 v_cndmask per iteration in a loop with 8 matrix instructions. Bitwise, −2 … −8 % on every shape this
kernel serves — which is why the serving gain below exceeds what these three boundaries alone predict.

4. Tile table. wqkv_a (1792×5120) had no row at all, so it fell back to a bf16 weight copy plus a
fake-quant pass; with a row it takes dot_scaled, and native_consumer_wants_fp8 flips so the fused norm emits
fp8 and the quantise launch disappears too: 38.4 → 32.1 µs at 768 rows. Also adds a 512 M-bucket, because
384 and 768 share bucket 1024 today and want different tiles — that alone is wo_b 22.5 → 17.2 µs at 384.
No extra memory: every shape here already carries a bf16 copy, since some bucket already resolves to
hipblaslt_bf16.

Measurements

MI355X / gfx950, layer 21 of DeepSeek-V4.1-Flash with real weights, TP4 shapes, --cache-regime prodlike
(weights and KV pool cold, as production sees them). Median of 100 samples of 4 calls in one HIP graph, two
processes per arm.

boundary, µs M=96 M=192 M=384 M=768
wqkv_a 18.3 → 18.3 22.1 → 22.2 27.2 → 26.3 38.4 → 32.1
wq_b 12.3 → 12.1 13.9 → 14.6 18.7 → 16.7 25.3 → 23.7
wo_b 15.0 → 14.0 15.4 → 15.3 22.5 → 17.2 26.2 → 24.4

Over 40 layers that is −0.33 ms per decode step at batch 64 and −0.39 at batch 128. The whole attention
block goes from 11.61 → 6.91 ms at batch 64 and 19.94 → 10.62 at batch 128.

On a server (TP4 + EP4 + DSpark, four alternating fresh servers, 12 drives of 256-token greedy requests per
arm): median wall per drive −4.6 % at batch 32 and −6.1 % at batch 64. One caveat worth knowing: the first
drive after a cold start pays a 30 s JIT compile of the four new tile configurations. That lands in server
warmup, not in serving.

Numerics

wo_b and wq_b are bitwise identical. wqkv_a is 2.3e-4 — half a bf16 ULP — because its route changes
from the bf16 copy to dot_scaled. Checked with a dump/compare harness over 12 shape × boundary pairs.

One thing that looks like a failure and is not: the boundary check reports an arity mismatch on input_norm
at M ≥ 384, because with a tile row in place the norm returns (fp8, ue8m0) instead of a grid-rounded bf16
tensor — that is the mechanism in change 4. Compared by value instead, dequantising the pair reproduces the
old bf16 exactly: 0 of 3,932,160 elements differ at M=768 K=5120, and at every other shape tested.

What is deliberately not here

The batch-aware KV split rule from the same campaign. That belongs in aiter, and is up as
ROCm/aiter#5736 against _decode_num_splits_occ instead.

Two decisions for you

  • The wo_a route is deterministic per shape but not bit-equal to the Triton tile it replaces, so any
    bitwise-reproducibility baseline taken on the old route needs re-taking.
  • Change 4 makes the fused norm emit fp8 at M ≥ 384, which changes that boundary's output type. Anything
    downstream that assumes a bf16 grid tensor there needs to accept the pair.

CI States

Latest PR Test (Base): ❌ Missing run-ci label -- add it to run CI tests.
Latest PR Test (Extra): ❌ Blocked -- run-ci is required first.
Latest PR Test (AMD ROCm 7.2): ➖ No AMD PR run found for this commit.

@tianxiaojiang4
tianxiaojiang4 marked this pull request as ready for review September 22, 2026 21:03
@tianxiaojiang4

Copy link
Copy Markdown
Author

@kevin-mii @HaiShaw @yichiche @chuyeh

Four measured gfx950 changes, stacked on dsv41-amd-main @ 7791f6c — the branch behind sgl-project#39857. The diff is 80 lines across 4 files; everything else in the range is Kevin's.

  • wo_a above the split-K ceiling. Past _SPLIT_K_MAX_M = 64 the 16x32 tile re-streams each group's whole weight once per 16 rows (~800 MB against a 16 MiB weight). A strided-batched bf16 GEMM written straight into the [T, G, R] layout the consumer wants: 125.3 -> 33.1 us at 768 rows.
  • Quantise once. Emit Mxfp8Activation rather than fp8-grid rounding, so wo_b stops re-quantising the same 2048-wide activation. Bitwise identical, -12% measured on the pair.
  • Skip the tile masks when the tile divides the problem. Loop-invariant; the disassembly shows 26 v_cndmask per iteration against 8 matrix instructions. Bitwise, -2 to -8%.
  • Tile table. wqkv_a (1792x5120) had no row and fell back to a bf16 copy plus a fake-quant pass; also adds a 512 M-bucket, since 384 and 768 share bucket 1024 today.

Why each of you, so you can ignore the parts that aren't yours:

@kevin-mii — these are your files. The actual question is routing, not code: would you rather I land these on dsv41-amd-main so they flow into sgl-project#39857, or should I hand you the four commits to fold in yourself? sgl-project#39857 is currently conflicting with main, so this PR's base goes stale on your next rebase either way.

@HaiShaw @yichiche @chuyeh — if you have bandwidth for one thing only, make it this: the route switch and the M-bucket table are both single-axis thresholds on M, and the crossover plausibly depends on weight footprint and cache residency too (the route comparison was measured cold; the end-to-end numbers are --cache-regime prodlike). sgl-project#39503 and sgl-project#39513 are the closest prior art I found for picking this kind of bound against CU share rather than one axis, which is why I'm tagging you. I'd rather have the threshold challenged now than fitted twice.

Related: sgl-project#39957 (merged 19 Sep) landed the SM100 fused inverse-RoPE + WO-A + MXFP8 path — same fusion idea as change 2, opposite row regime (<=32 vs >64). I'd like to confirm the Mxfp8Activation contract here matches what landed there.

Apologies for the venue — this targets Kevin's fork because the four files aren't on sgl-project:main yet. Happy to re-open upstream if that's easier to review.

… fixes

Four measured changes to the gfx950 decode projections, on top of dsv41-amd-main.
All are gated against the unmodified branch and validated on a live server.

1. wo_a above the split-K ceiling. Past _SPLIT_K_MAX_M the 16x32 grid tile
   re-streams each group's whole weight once per 16 rows, ~800 MB of traffic for
   a 16 MiB weight. A strided-batched bf16 GEMM written straight into the
   [T, G, R] layout the consumer wants, with the same fp8-grid rounding applied
   once afterwards, takes it from 125.3 to 33.1 us at 768 rows.
   SGLANG_OPT_WO_A_LARGE_M_ROUTE=triton restores the old tile.

2. Quantise once. On that route the result is rounded onto the fp8 grid in its
   own pass, and wo_b - seeing no input scale - then quantises the same 2048-wide
   activation again. Emitting Mxfp8Activation instead removes one of the two
   passes and is bitwise identical. The pair drops 12% at every M >= 96.
   SGLANG_OPT_WO_A_LARGE_M_EMIT=grid restores the current emission.

3. Skip the tile masks when the tile divides the problem. They are loop-invariant
   and cost 26 v_cndmask per iteration in a loop with 8 matrix instructions.
   Bitwise, and worth 2-8% on every shape this kernel serves.

4. Tile table. wqkv_a (1792x5120) had no row at all, so it fell back to a bf16
   weight copy plus a fake-quant pass; with a row it takes dot_scaled and the
   fused norm starts emitting fp8, which deletes the quantise launch too.
   Adds a 512 M-bucket, because 384 and 768 rows share bucket 1024 today and want
   different tiles. No extra memory: every shape here already carries a bf16 copy.

Per decode step over 40 layers: -0.33 ms at batch 64, -0.39 at batch 128. The
attention block goes from 11.61 to 6.91 ms at batch 64 and 19.94 to 10.62 at 128.
On a TP4 + EP4 + DSpark server, median wall per 256-token drive falls 4.6% at
batch 32 and 6.1% at batch 64.

Numerics: wo_b and wq_b bitwise identical; wqkv_a 2.3e-4 (half a bf16 ULP) since
its route changes. The norm's output type changes at M >= 384 - codes and scales
instead of grid-rounded bf16, which is the mechanism - and dequantising
reproduces the old value bitwise at every shape tested.

Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
@tianxiaojiang4
tianxiaojiang4 force-pushed the amd-gfx950-decode-projections branch from 050c2fc to 1461076 Compare September 23, 2026 17:52
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant