Repository navigation
[Triton/Gluon] [JIT] [Feature] Add tuned native group32 A8W8 Triton GEMM for gfx950 - #5750
Conversation
Integrate unshuffled E4M3/E8M0 inputs into gemm_a8w8_blockscale. Keep microscaling and packed-K kernels in Triton, with gfx950 tuning tables loaded by get_gemm_config. Extend shared loading with JSON M bounds and reduce config-copy overhead while retaining defensive nested copies. Reuse the shared non-atomic split-K reduction and preserve legacy routes. Validated with 98 AITER tests, 73,728 matching config selections, and 285 bitwise-equal BF16 migration cases; repeated 40 GPU cases six times. Signed-off-by: Lingpeng Jin <103567126+valarLip@users.noreply.github.com>
Sweep 18 gfx950 geometries across TP4, TP2, and attention-DP local projections. Store tile, split-K, launch, and M-bound selections in the existing JSON family; retain baseline intervals without stable gains. Keep the existing large-M fallbacks and merge equal adjacent configs. Exercise N-first masking independently of the selected traversal. Validation: 49,304 candidates screened, 603 regression cases, 98 AITER tests and 85 ATOM tests passed. Paired regression geometric-mean speedup is 1.179x; changed configurations have no sampled slowdown above 1%. Signed-off-by: Lingpeng Jin <103567126+valarLip@users.noreply.github.com>
Signed-off-by: Lingpeng Jin <103567126+valarLip@users.noreply.github.com>
Signed-off-by: Lingpeng Jin <103567126+valarLip@users.noreply.github.com>
🏷️ CI GuideRuns automatically on every PR:
Extended tests (opt-in via labels):
PR title tags & labels: |
Signed-off-by: Lingpeng Jin <103567126+valarLip@users.noreply.github.com>
There was a problem hiding this comment.
Copilot review overview
🟡 Changes recommended
One or more issues must be addressed before approval.
Get a fresh assessment by requesting another Copilot review.
Review effort: Lite
Findings: 4
Open (6)
Avoid ambiguous group_n inference for N <= 32 · New Detect raw uint8 E8M0 scale views before legacy dispatch · New Include all tuned launch knobs in kernel representations · New Reinterpret FP8 operands as uint8 before tl.dot_scaled · New Test production config resolution without overriding config · New Add benchmark coverage for native group32 kernels · New
What changed in this PR
Adds a gfx950 native group32 E4M3/E8M0 Triton GEMM backend with configuration-driven dispatch, tuning tables, and validation tests.
Changes:
- Adds native packed and split-K Triton kernels with public dispatch support.
- Extends GEMM configuration loading with JSON-defined M bounds.
- Adds gfx950 tuning data, CSV backend selection, and tests.
| File | Description |
|---|---|
| op_tests/triton_tests/gemm/basic/test_gemm_a8w8_blockscale_group32.py | Updated as part of this pull request. |
| op_tests/test_gemm_group32_interface.py | Updated as part of this pull request. |
| aiter/ops/triton/utils/gemm_config_utils.py | Updated as part of this pull request. |
| aiter/ops/triton/gemm/basic/gemm_a8w8_blockscale_group32.py | Updated as part of this pull request. |
| aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a8w8_blockscale_group32/GEMM-A8W8_BLOCKSCALE_GROUP32-N=8192-K=1280.json | Updated as part of this pull request. |
| aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a8w8_blockscale_group32/GEMM-A8W8_BLOCKSCALE_GROUP32-N=5120-K=8192.json | Updated as part of this pull request. |
| aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a8w8_blockscale_group32/GEMM-A8W8_BLOCKSCALE_GROUP32-N=5120-K=576.json | Updated as part of this pull request. |
| aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a8w8_blockscale_group32/GEMM-A8W8_BLOCKSCALE_GROUP32-N=5120-K=4096.json | Updated as part of this pull request. |
| aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a8w8_blockscale_group32/GEMM-A8W8_BLOCKSCALE_GROUP32-N=5120-K=2304.json | Updated as part of this pull request. |
| aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a8w8_blockscale_group32/GEMM-A8W8_BLOCKSCALE_GROUP32-N=5120-K=2048.json | Updated as part of this pull request. |
| aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a8w8_blockscale_group32/GEMM-A8W8_BLOCKSCALE_GROUP32-N=5120-K=15360.json | Updated as part of this pull request. |
| aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a8w8_blockscale_group32/GEMM-A8W8_BLOCKSCALE_GROUP32-N=5120-K=1152.json | Updated as part of this pull request. |
| aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a8w8_blockscale_group32/GEMM-A8W8_BLOCKSCALE_GROUP32-N=4608-K=5120.json | Updated as part of this pull request. |
| aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a8w8_blockscale_group32/GEMM-A8W8_BLOCKSCALE_GROUP32-N=4096-K=1280.json | Updated as part of this pull request. |
| aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a8w8_blockscale_group32/GEMM-A8W8_BLOCKSCALE_GROUP32-N=32768-K=1280.json | Updated as part of this pull request. |
| aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a8w8_blockscale_group32/GEMM-A8W8_BLOCKSCALE_GROUP32-N=25600-K=6144.json | Updated as part of this pull request. |
| aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a8w8_blockscale_group32/GEMM-A8W8_BLOCKSCALE_GROUP32-N=2304-K=5120.json | Updated as part of this pull request. |
| aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a8w8_blockscale_group32/GEMM-A8W8_BLOCKSCALE_GROUP32-N=1792-K=5120.json | Updated as part of this pull request. |
| aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a8w8_blockscale_group32/GEMM-A8W8_BLOCKSCALE_GROUP32-N=16384-K=1280.json | Updated as part of this pull request. |
| aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a8w8_blockscale_group32/GEMM-A8W8_BLOCKSCALE_GROUP32-N=1536-K=5120.json | Updated as part of this pull request. |
| aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a8w8_blockscale_group32/GEMM-A8W8_BLOCKSCALE_GROUP32-N=1152-K=5120.json | Updated as part of this pull request. |
| aiter/ops/triton/configs/gfx950/triton/gemm/gemm_a8w8_blockscale_group32/DEFAULT.json | Updated as part of this pull request. |
| aiter/ops/triton/configs/CLAUDE.md | Updated as part of this pull request. |
| aiter/ops/triton/_triton_kernels/gemm/basic/gemm_a8w8_blockscale_group32.py | Updated as part of this pull request. |
| aiter/ops/gemm_op_a8w8.py | Updated as part of this pull request. |
| aiter/jit/core.py | Updated as part of this pull request. |
| aiter/configs/a8w8_blockscale_group32_untuned_gemm.csv | Updated as part of this pull request. |
| aiter/configs/a8w8_blockscale_group32_tuned_gemm.csv | Updated as part of this pull request. |
💡 Configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
| ) | ||
|
|
||
| assert WQ.ndim == w_scale.ndim == 2, "Expected matrix weights and scales" | ||
| group_n = 32 if w_scale.shape[0] == -(-WQ.shape[0] // 32) else 1 |
| is_group32 = ( | ||
| x_scale.dtype == dtypes.fp8_e8m0 | ||
| and w_scale.dtype == dtypes.fp8_e8m0 | ||
| and x_scale.ndim == 2 | ||
| and XQ.ndim == 2 | ||
| and x_scale.shape[1] == XQ.shape[1] // 32 |
| _gemm_group32_repr = make_kernel_repr( | ||
| "_gemm_a8w8_blockscale_group32_kernel", | ||
| [ | ||
| "BLOCK_SIZE_M", | ||
| "BLOCK_SIZE_N", | ||
| "BLOCK_SIZE_K", | ||
| "GROUP_N", | ||
| "SPLITK_BLOCK_SIZE", | ||
| "N_FIRST", | ||
| ], | ||
| ) | ||
| _gemm_group32_packed_repr = make_kernel_repr( | ||
| "_gemm_a8w8_blockscale_group32_packed_kernel", | ||
| ["BLOCK_SIZE_M", "BLOCK_SIZE_N", "BLOCK_SIZE_K", "K_PACK"], | ||
| ) |
| x = x.view(torch.float8_e4m3fn) if x.dtype == torch.uint8 else x | ||
| w = w.view(torch.float8_e4m3fn) if w.dtype == torch.uint8 else w |
| config, tuned = get_gemm_config("GEMM-A8W8_BLOCKSCALE_GROUP32", m, 8192, 1280) | ||
| assert tuned | ||
| if packed: | ||
| assert "packed" in config | ||
| else: | ||
| # Exercise N-first masking independently of which traversal wins tuning. | ||
| config.pop("packed", None) | ||
| config["N_FIRST"] = True | ||
| x, w, xs, ws = _operands(m, 8193, 1312) | ||
| actual = gemm_a8w8_blockscale_group32(x, w, xs, ws, dtype=dtype, config=config) |
| def gemm_a8w8_blockscale_group32( | ||
| x: torch.Tensor, | ||
| w: torch.Tensor, | ||
| x_scale: torch.Tensor, | ||
| w_scale: torch.Tensor, | ||
| dtype: torch.dtype = torch.bfloat16, | ||
| y: torch.Tensor | None = None, | ||
| weight_group_rows: int = 32, | ||
| split_k: int | None = None, | ||
| config: dict | None = None, |
Small-M group32 GEMMs on narrow or deep projections are starved for CTAs: N=512 gives the packed kernel 32 CTAs on 256 CUs, and splitting K costs a second reduce launch. Both kernels now take FUSED_SPLITS, which sums the split partials in the same launch. A 1D grid places every split of a tile on one XCD (CTAs go to XCDs by pid % 8, a multiple of every gfx950 XCD count). That XCD's L2 is then the splits' coherence point, so the last CTA to arrive can read the other partials after a workgroup-scope acq_rel counter add: no agent-scope release, and no buffer_wbl2 of the whole L2. With agent scope the same fused kernels are 2.3x slower, which is why a plain fused split-K lost to the two-launch path. Partials are summed in split order whatever the arrival order, so results are deterministic, and the last CTA re-zeroes its counter, so graph replay is safe. Counters follow the a16w16 split-K semaphores: one zeroed buffer per (device, stream) from persistent_alloc, so concurrent streams never share one. Where a stream's first use falls inside a capture that cannot allocate outside the graph pool (torch < 2.10), the launch takes the unfused path instead of failing the capture. The fused path is selected only by the tuning table (FUSED_SPLITK on a tile config, NUM_KSPLIT on a packed config). An explicit split_k keeps the existing reduce-kernel path. With the original configs, all 567 sampled outputs across the 18 V4.1 geometries, tails, FP32/BF16 and explicit split_k are bitwise identical to before. Tests cover the reference match for both variants and output dtypes, 1x32 weight scales, 50-call determinism with counters re-zeroed, graph replay on a warmed stream, first use inside a capture, and interleaved launches on two streams.
57 M intervals of 11 V4.1 geometries now use the fused split-K path, 42 at M <= 64 and 15 above. Candidates were the best fused tile and packed configs from full sweeps; an interval changed only if a candidate was at least 3% faster at both of its bounds. Every changed interval was then re-measured against its previous config in the same process (run_perftest, rotated inputs, median of 5): at each CSV M inside it, or at its lower bound, midpoint and upper bound where no CSV M falls inside. Three intervals that did not hold up were left unchanged. No measured point is slower than the previous config. Speedup over the previous config, geometric mean of the rows re-measured: o_b.dp M <= 64 1.26x draft_kvf M 129..320 1.22-1.26x draft_kv M <= 64 1.22x sh_up.tp4 M 321..448 1.26x sh_up.tp4 M <= 64 1.19x sh_up.tp2 M 129..192 1.20x draft_kvf M <= 64 1.13x o_b.dp M 129..192 1.20x o_b.tp2 M <= 64 1.12x draft_kv M 513..576 1.14x The V4.1 model CSV previously left 188 us fields and every tflops, bw and errRatio field empty. All 450 rows are now timed in one run with the final configs, through gemm_a8w8_blockscale, with run_perftest and rotated inputs, so the us column comes from a single method. tflops and bw use the a8w8 blockscale tuner formula (bw excludes scales). errRatio is checkAllclose at its default tolerance against an FP64 reference, as in the tuner: at most 8.7e-4, and 0 rows above 0.05.
…une for group32 GEMM Retune all 18 V4.1 geometries on the repo's 68 DSv4 M buckets (fused split-K, packed K_PACK=1, .cg weight loads); drop unused GROUP_SIZE_M; re-time the 1224-row CSV.
|
For the record: the in-launch split-K reduction on one XCD was proposed here in d63b959 (pushed 2026-09-24 16:48 UTC) and adopted by vllm-project/vllm#58510 in c3c00323dd (pushed 19:23 UTC the same day). That PR's packed small-M kernel also follows the one this PR added in 4ffc0a7 (2026-09-22). vllm-project/vllm#58659 adds the attribution. |
Depends on ROCm/aiter#5750, which adds a native MXFP8 GEMM (e4m3 with e8m0 scales per 32 along K, 32x32-block weight scales) behind aiter.gemm_a8w8_blockscale. With --fp8-gemm-backend aiter, the DSv4.1 fp8 dense linears on gfx950 run on that kernel instead of the sglang native MXFP8 route (mxfp8_gemv / dot_scaled / hipBLASLt bf16).
Depends on ROCm/aiter#5750, which adds a native MXFP8 GEMM (e4m3 with e8m0 scales per 32 along K, 32x32-block weight scales) behind aiter.gemm_a8w8_blockscale. With --fp8-gemm-backend aiter, the DSv4.1 fp8 dense linears on gfx950 run on that kernel instead of the sglang native MXFP8 route (mxfp8_gemv / dot_scaled / hipBLASLt bf16).
Depends on ROCm/aiter#5750, which adds a native MXFP8 GEMM (e4m3 with e8m0 scales per 32 along K, 32x32-block weight scales) behind aiter.gemm_a8w8_blockscale. With --fp8-gemm-backend aiter, the DSv4.1 fp8 dense linears on gfx950 run on that kernel instead of the sglang native MXFP8 route (mxfp8_gemv / dot_scaled / hipBLASLt bf16).
Depends on ROCm/aiter#5750, which adds a native MXFP8 GEMM (e4m3 with e8m0 scales per 32 along K, 32x32-block weight scales) behind aiter.gemm_a8w8_blockscale. With --fp8-gemm-backend aiter, the DSv4.1 fp8 dense linears on gfx950 run on that kernel instead of the sglang native MXFP8 route (mxfp8_gemv / dot_scaled / hipBLASLt bf16).
Depends on ROCm/aiter#5750, which adds a native MXFP8 GEMM (e4m3 with e8m0 scales per 32 along K, 32x32-block weight scales) behind aiter.gemm_a8w8_blockscale. With --fp8-gemm-backend aiter, the DSv4.1 fp8 dense linears on gfx950 run on that kernel instead of the sglang native MXFP8 route (mxfp8_gemv / dot_scaled / hipBLASLt bf16).
Depends on ROCm/aiter#5750, which adds a native MXFP8 GEMM (e4m3 with e8m0 scales per 32 along K, 32x32-block weight scales) behind aiter.gemm_a8w8_blockscale. With --fp8-gemm-backend aiter, the DSv4.1 fp8 dense linears on gfx950 run on that kernel instead of the sglang native MXFP8 route (mxfp8_gemv / dot_scaled / hipBLASLt bf16).



DeepSeek-V4.1-Flash projections use unshuffled E4M3 weights and compact E8M0 group32 scales. This adds a gfx950 Triton backend for those operands to
gemm_a8w8_blockscale, together with tuned configurations for TP4, TP2 and attention-DP local GEMMs.Implementation and dispatch
split_kpartition count, preserving explicit caller controls while keeping configured backend selection first.get_CKGEMM_configbefore selecting a fallback. Native E8M0 operands have a separateAITER_CONFIG_GEMM_A8W8_BLOCKSCALE_GROUP32CSV family, so an identical M/N/K in the FP32 128x128 table cannot select an incompatible kernel. Unsupported configured backends raise instead of silently falling back.libtype=triton. Tile, split-K and packed-variant settings remain in the existing Triton JSON configuration system. Of these rows, 262 reuse exact-shape measured timings; remaining timing fields are empty.M_BOUNDS, retaining explicit-argument precedence and legacy bounds. Config copies preserve nested mutation isolation while avoiding per-scalardeepcopyoverhead.The new native backend requires gfx950. Existing FP32 128x128 operands retain their CK/CKTile paths. Group32 currently supports the Triton backend; configuring legacy CK/CKTile for its E8M0 operands is rejected.
Performance
Hardware: AMD gfx950, 256 CUs; PyTorch 2.9.1 ROCm 7.1.1 and Triton 3.8.0 ROCm build. Measurements use BF16 output, CUDA graphs and at least 512 MiB of rotating weights. The baseline is the initial imported native group32 configuration in the first commit of this PR.
The sweep evaluated 49,304 candidates across 234 representative M/N/K cases and validated 603 cases, including interval boundaries. The geometric-mean speedup is 1.179x (15.18% lower GEMM latency). Of 300 intervals checked, 252 use improved configurations and 48 retain the baseline. The largest sampled slowdown among changed configurations is 0.97%. Measurements affected by other GPU users were discarded and rerun on available cards.
These are standalone GEMM results. End-to-end serving throughput was not measured. Tuning covers M through 4096; larger M retains the original JSON fallback. Model CSV entries through M=16384 choose the backend without claiming newly tuned tiles at those larger sizes.
Validation
git diff --checkpasses.The sweep compares sampled outputs against an independent FP64 reference and complete BF16 outputs against the baseline. BF16 outputs are bitwise equal in 384/603 cases; the maximum BF16 relative RMS difference is 1.052e-4. One pre-existing FP32 reference-boundary outlier at M=97, N=5120, K=1152 remains: maximum absolute error changes from 0.09451294 to 0.09445190. The experiment accepts that point only when candidate error does not exceed baseline error; repository test tolerances are unchanged.