Repository navigation
[Triton/Gluon] [Config] Add gfx950 gemm_a16w16 config for N=288, K=4096 - #5933
Conversation
GLM-5.3-Flash routes over 288 experts with hidden size 4096, so each of its 42 MoE layers runs a bf16 N=288, K=4096 gate GEMM. gfx950 has no per-shape file for it and falls back to DEFAULT.json, whose M_LEQ_8 and M_LEQ_32 entries set waves_per_eu=8. That caps the kernel at 64 VGPRs; it spills 16 (M<=8) and 21 (M<=32) registers and takes 27-35 us per call where a spill-free kernel takes ~5 us. The new file drops the occupancy hint (waves_per_eu=0) and the .cg load modifier for M<=64 and keeps a single kernel (NUM_KSPLIT=1): a BM=8/16, BN=16, BK=512 tile, 90-98 VGPRs, no spills. The entries were picked from an 8,640-config NUM_KSPLIT=1 sweep timed in a HIP graph over 42 rotating weights, then checked with bench_gemm_a16w16.py. M_LEQ_128 and above are copied verbatim from DEFAULT.json, so larger M resolves exactly as before.
🏷️ CI GuideRuns automatically on every PR:
Extended tests (opt-in via labels):
One backend per PR: PR title tags & labels: |
There was a problem hiding this comment.
Copilot review overview
🟢 Approval recommended
The configuration is valid, complete, correctly located, and preserves the documented defaults.
Review effort: Balanced
Findings: None
What changed in this PR
Adds a gfx950 Triton tuning profile for GLM-5.3-Flash’s BF16 router GEMM.
Changes:
- Tunes M ≤ 64 for N=288, K=4096.
- Preserves default configurations for larger M values.
| File | Description |
|---|---|
aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a16w16/GEMM-A16W16-N=288-K=4096.json |
Adds shape-specific GEMM tuning tiers. |
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
Motivation
GLM-5.3-Flash's MoE router (288 experts over hidden size 4096) runs a bf16 GEMM of N=288, K=4096 in each of its 42 MoE layers, and at decode batch 1–8 each call takes 27 us where a spill-free kernel takes 5 us. gfx950
gemm_a16w16has no file for this shape, so it falls back toDEFAULT.json, whoseM_LEQ_8andM_LEQ_32entries setwaves_per_eu=8. That caps the kernel at 64 VGPRs and it spills 16–21 registers. The newGEMM-A16W16-N=288-K=4096.jsondrops that hint and.cgfor M ≤ 64 and keeps one kernel (no split-K).M_LEQ_128and above are copied verbatim fromDEFAULT.json.Test Plan
MI355X (gfx950), ROCm 10.0.0, Triton 3.8.0,
rocm/sgl-dev:v0.5.20-rocm10-mi35x-20260928.Step 1 — a
NUM_KSPLIT=1sweep picks the entries. 8,640 configs over M=1–64, timed in a HIP graph over 42 rotating weights (one per MoE layer, so the weight read misses L2). Among configs within 1% of the fastest, the one withwaves_per_eu=0was kept.Step 2 —
bench_gemm_a16w16.pyA/Bs the file on upstreammain, moving the JSON in and out ofaiter/ops/triton/configs/gfx950/triton/gemm/gemm_a16w16/. Median of 3 runs per point:python3 op_tests/op_benchmarks/triton/bench_gemm_a16w16.py --shape 288 4096 --metric time
Correctness:
gemm_a16w16vstorch.nn.functional.linear, M ∈ {1…256}, layouts TN/TT/NN/NT, with and without bias — 128 cases, 0 failures. End to end: SGLangbench_serving, TP4, same JSON in/out of the image's config dir.Test Result
Kernel, us/call. "bench" is the Step 2 command; "graph" is the Step 1 harness with 42 rotating weights.
M_LEQ_4M_LEQ_4M_LEQ_4M_LEQ_8M_LEQ_16M_LEQ_32M_LEQ_64M_LEQ_128(unchanged)M_LEQ_256(unchanged)No M is slower. M=16 is flat on the bench because a single hot weight hides DEFAULT's cost there; with rotating weights it is 7.12 → 5.10 us. M=128/256 resolve to the same entries as before, field for field.
End to end, GLM-5.3-Flash-Quark-MXFP4, TP4, input 8192 / output 1024:
In SGLang, decode batches above 16 route the router GEMM through
aiter.tuned_gemmto hipBLASLt, so the gain is at concurrency ≤ 16; theM_LEQ_32/M_LEQ_64entries matter only for callers that reachgemm_a16w16directly.