Repository navigation
[AMD][gfx950] wo_a route above 64 rows, quantise once, and MXFP8 tile fixes - #10
tianxiaojiang4 wants to merge 1 commit into
Conversation
|
@kevin-mii @HaiShaw @yichiche @chuyeh Four measured gfx950 changes, stacked on
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 @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 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 Apologies for the venue — this targets Kevin's fork because the four files aren't on |
7791f6c to
5a3ee1f
Compare
… 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>
050c2fc to
1461076
Compare
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#39857instead if you would rather carry them.
The changes
1.
wo_aabove the split-K ceiling. Past_SPLIT_K_MAX_Mthe 16×32 grid tile re-streams each group's wholeweight 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 onceafterwards, takes it from 125.3 → 33.1 µs at 768 rows.
SGLANG_OPT_WO_A_LARGE_M_ROUTE=tritonrestores theold 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— whichsees no input scale — then quantises the same 2048-wide activation again. Emitting
Mxfp8Activationinsteadremoves one of the two passes and is bitwise identical;
wo_b's own quantise kernel disappears because theoperand 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=gridrestores the current emission.3. Skip the tile masks when the tile divides the problem. They are loop-invariant, and the disassembly shows
26
v_cndmaskper iteration in a loop with 8 matrix instructions. Bitwise, −2 … −8 % on every shape thiskernel 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 afake-quant pass; with a row it takes
dot_scaled, andnative_consumer_wants_fp8flips so the fused norm emitsfp8 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_b22.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.
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_bandwq_bare bitwise identical.wqkv_ais 2.3e-4 — half a bf16 ULP — because its route changesfrom 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_normat M ≥ 384, because with a tile row in place the norm returns
(fp8, ue8m0)instead of a grid-rounded bf16tensor — 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_occinstead.Two decisions for you
wo_aroute is deterministic per shape but not bit-equal to the Triton tile it replaces, so anybitwise-reproducibility baseline taken on the old route needs re-taking.
downstream that assumes a bf16 grid tensor there needs to accept the pair.
CI States
Latest PR Test (Base): ❌ Missing
run-cilabel -- add it to run CI tests.Latest PR Test (Extra): ❌ Blocked --
run-ciis required first.Latest PR Test (AMD ROCm 7.2): ➖ No AMD PR run found for this commit.