Skip to content

feat(cco): gemm_ar MXFP8 support, and the GEMM usable without a collective - #690

Merged
yangyuhuiling merged 8 commits into
mainfrom
xiangch/gemm-ar-dsv41-shapes
Sep 21, 2026
Merged

yangyuhuiling merged 8 commits into
mainfrom
xiangch/gemm-ar-dsv41-shapes

Conversation

@yangyuhuiling

Copy link
Copy Markdown
Contributor

What this adds

mori.ops.gemm_ar gains support for 32-wide ue8m0 (MXFP8), the quantisation
DeepSeek-V4.1-Flash actually ships (weight_block_size [32, 32],
scale_fmt "ue8m0"), and two ways to use the GEMM without a collective attached.

op for status
GemmAllReduceOp a RowParallelLinear at prefill M existing, now serves mxfp8 as well as blockscale
Mxfp8GemmOp (gemm.py) any mxfp8 linear at prefill M new
Mxfp8GemvOp (gemv.py) any mxfp8 linear at decode M (<= 32) new

The second and third exist because at V4.1-Flash's shapes most of the win is
the multiply, not the overlap
-- fusing is worth 11-15% at best there, where
the GEMM alone is worth ~30% against the route SGLang would otherwise take. And
a ColumnParallelLinear has no all-reduce to fuse with at all.

CDNA4's v_mfma_scale_f32_16x16x128_f8f6f4 takes the ue8m0 scales as
instruction operands, so the 32-wide form needs no dequantisation arithmetic
at all -- unlike the existing _BlockScaleK promote/rescale chain, which is
there precisely because a 128-wide block cannot be expressed that way.

Headline numbers

MI355X (gfx950). Against mxfp8_native_blockscaled_linear, identical operands,
bf16 in and bf16 out with quantisation included on both sides, cold:

wq_b 8192x1280 wo_b 5120x2048
M=2048 -24% -34%
M=16384 -43% -27%

Fused GEMM + all-reduce at V4.1-Flash's wo_b, TP4, M=16384: 1192us against
1592 for split/bf16, -25.1%
on fused-sdma + fp8 gather over LSA.

Two layout results came out of making the GEMM fast, both about addresses
rather than bytes, and both are documented with the hardware counters that
settle them:

  • Scales must be K-block major. Same instruction count, 3.2x the cache
    accesses row-major, +50-67% on the GEMM.
  • Four M tiles pack into one scale dword, via opsel_b on the scaled MFMA.
    -29.1% SQ_INSTS_VMEM with VALU flat, worth -5 to -7.5%.

Correctness

  • 66 single-GPU cases in tests/python/cco/test_gemm_ar_op.py, plus the
    existing 2-rank worker cases.
  • Where the reference takes tl.dot_scaled, mori agrees with it bit for
    bit
    -- same instruction, same operands, and ue8m0 scales are exact powers of
    two. That is asserted, not just tolerated.
  • The GEMV is bit-identical to SGLang's mxfp8_gemv at every M and both shapes.
  • The 112-config correctness grid (check_gemv_grid.sh) passes on both real
    shapes and every M in {1, 2, 8, 16, 32}.

Notes for review

_gemm_a8w8_8wave.py is vendored from aiter and now diverges in two places,
not one.
Mfma16x16x128 grew the scaled form of the atom: opsel_b_per_tile
on the constructor, scale_a/scale_b on call. Both default to off and the
unscaled path is byte-for-byte what it was, so a re-vendor stays a merge. The
file's header says so.

_PinnedLaunch reaches into FlyDSL internals (_call_state_cache, _sig)
to skip per-dispatch bookkeeping that measured 181.6us of Python per fused call
against 5.9us of actual hipModuleLaunchKernel. It verifies the argument tuple
element-by-element before pinning and falls back to the normal path on any
surprise, so a FlyDSL change degrades performance rather than correctness. It
is the difference between a layer being host-bound and not.

Two quantisation modes exist only to be measured against: mxfp8_unpacked
and mxfp8_row are the scale layouts mxfp8 was chosen over. They are ~5 lines
of plumbing and they keep the choice reproducible; happy to drop them if you
would rather not carry them.

bench_gemm_ar.py keeps its single-call graph timer deliberately. On this
box that carries a ~13.4us replay floor, which is 1-4% of the 300-1600us fused
all-reduce it measures and identical across the modes being compared; changing
it would move every number already published from it. New benchmarks of single
small kernels use benchmark/cco/flydsl/gemm_ar/timing.py instead, which
amortises the capture and reports cold as well as hot. The reasoning is in the
README under "Measurement traps", along with the two ways this branch got it
wrong before getting it right.

Docs moved. docs/MORI-GEMM-AR-BENCHMARK.md is now how-to-run plus a
summary (744 -> 177 lines); every measurement table lives in
python/mori/ops/gemm_ar/README.md, next to the code it is about.

Not in scope

  • scatter_dtype='fp8' still raises NotImplementedError. Its regions are
    sized but nothing writes the payload; that leg is the largest remaining lever
    on the fused path.
  • The shared expert's gate_up (N=1152) is not expressible -- N is 4.5 tiles of
    256. Measured on the narrow tile anyway and it never wins, so supports_gemm
    is left refusing it.

@yangyuhuiling
yangyuhuiling force-pushed the xiangch/gemm-ar-dsv41-shapes branch from 1a2d8fd to 5e6ca07 Compare September 20, 2026 05:34
xiangch and others added 6 commits September 20, 2026 08:53
…ships

`GemmAllReduceOp` served 1x128 / 128x128 fp8 block scales. DeepSeek-V4.1-Flash
quantises 32-wide ue8m0 on both operands (`weight_block_size [32, 32]`,
`scale_fmt "ue8m0"`), which is a different operand contract and a different
kernel, not a parameter of the old one.

CDNA4's `v_mfma_scale_f32_16x16x128_f8f6f4` takes the ue8m0 scales as
*instruction operands*, so this form needs no dequantisation arithmetic at all
-- no second accumulator, no running rescale. `_BlockScaleK`'s promote/rescale
chain exists precisely because a 128-wide block cannot be expressed that way.

Two layout results came out of making it fast, both about addresses rather than
bytes, and both carry the counters that settle them:

* **Scales must be K-block major.** A block group's sixteen lanes want sixteen
  consecutive rows of one K block. K-block major coalesces them; the quantiser's
  own row-major `[M, K/32]` spreads them K/32 bytes apart -- same instruction
  count, 3.2x the cache accesses, +50-67% on the GEMM.
* **Four M tiles pack into one scale dword.** `opsel_b` on the scaled MFMA is an
  atom-time attribute naming which byte of a 32-bit scale operand to read, so
  one load serves four tiles and the select is free. -29.1% `SQ_INSTS_VMEM` with
  VALU flat, worth -5 to -7.5%. A quantiser can emit the layout directly, at no
  extra traffic.

`_PinnedLaunch` resolves each compiled kernel's dispatch once. `JitFunction`
re-derived its cache key per call -- an inspect bind, a 35-global snapshot, a
drift check -- which measured 181.6us of Python per fused call against 5.9us of
actual `hipModuleLaunchKernel`, enough to make the layer host-bound. It verifies
the argument tuple element by element before pinning and falls back to the
normal path on any surprise, so a FlyDSL change costs performance, not
correctness.

The fp8 all-gather leg gets its two knobs settled by measurement: pull it over
LSA rather than pushing over SDMA (2.0-4.0%), and do not fuse the quantise into
the reduce (1.2-1.9%). Both hold at every M and in all three modes.

`_gemm_a8w8_8wave.py` is vendored from aiter and now diverges in two places
rather than one: `Mfma16x16x128` grew the scaled form of the atom. Both new
arguments default to off and the unscaled path is byte-for-byte what it was, so
a re-vendor stays a merge. The file's header says so.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Most of what this integration is worth turns out to be the multiply, not the
overlap. At V4.1-Flash's shapes the GEMM is ~180us against 1130-1410us of
communication, so hiding it entirely would be worth 11-18% and fusing collects
about half of that -- while the GEMM alone is worth ~30% against the route
SGLang would otherwise take. And a `ColumnParallelLinear` has no all-reduce to
fuse with at all, so `wq_b` was unreachable.

`Mxfp8GemmOp` (`gemm.py`) is the same kernel with the epilogue tail compiled
out: no communicator, no window, an ordinary tensor out. Two things about it are
not obvious. M only has to be a multiple of 64, not of BLOCK_M, because the grid
is ceildiv and the tail block masks -- 64 is the packed scale's group. And there
is exactly one compile per (N, K) per tile width, because `c_m` is a runtime
argument; a per-M cache looks harmless until a server drives it, where every
prefill batch pays a fresh multi-second FlyDSL compile.

Which N tile to use is decided by the resulting **grid**, not by M. The tile
exists to double the grid when the wide one is launch-starved, and the grid is
`ceildiv(M,256) * (N/256)`, so a threshold on M alone is only ever right for the
N it was fitted to. Across 110 measured points the separation is total: below
128 workgroups the 128-wide tile wins every single time by 9-17%, from 129 the
256-wide one wins every single time by 26-29%. `_WIDE_TILE_MIN_GRID = 140` sits
in the gap at zero regret.

`Mxfp8GemvOp` (`gemv.py`, `kernels_gemv.py`) is a different kernel for decode's
token counts, where a 256-row tile has nothing to fill it. It inverts the same
scaled MFMA -- the weight is the A operand and the tokens are B, because the
weight is what there is a lot of -- streams the weight from global with no LDS
staging, and splits K across a workgroup's waves, reducing through LDS in a
fixed order so a row stays batch-invariant. Against SGLang's `mxfp8_gemv` on the
same fp8 bytes it is 3-8% faster on the two tuned shapes and **bit-identical**
at every M, which is the check that matters: same instruction, same operands, so
a layout disagreement shows up as a mismatch rather than a tolerance.

Its margin is tuning, and that is recorded next to the table: shapes falling
back to the heuristic land inside +-2%, except `wo_a`, which loses 12-16%.

The operands are deliberately not uniform. Both ops take `preshuffle_b`'s
weight, so a server shuffles once, but the GEMM wants the A scale through
`preshuffle_a_scale` while the GEMV takes both scales exactly as the checkpoint
stores them. That is not an inconsistency: the GEMM's sixteen lanes read sixteen
rows of one K block and have to coalesce, where the GEMV's are sixteen tokens,
M is at most 32, and the whole A scale is under a kilobyte.

Tests cover both ops at both real shapes, the ragged and unpadded M contracts,
and the GEMV across its configuration space -- each configuration in its own
process, because a FlyDSL compile failure takes the interpreter with it. Every
masked address in the GEMV is pushed out of its buffer's records or clamped:
0xFF is NaN in both ue8m0 and e4m3, and NaN times a zeroed weight is still NaN.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
…st most

The measurements, and the harness that has to be right for them to mean
anything. Most of the work here was finding out that it was not.

**Three measurement traps, each of which produced a confident wrong number.**
A single-call CUDA-graph replay has a ~13.4us floor on this box, which is most
of any small-M measurement -- `timing.py` amortises the capture. A repeated call
reads the weight out of the 256MB LLC at 1.7x the bandwidth a forward pass gets,
so cold is the number that predicts a server -- it rotates copies past the
cache. And getting *that* right took two more corrections: the ring advances at
capture time, so a graph of `reps` calls only ever touches `reps` copies; and
the ring has to be sized so each tensor clears the cache on its own, not so
their sum does, because a caller may hand over several weights of which the
kernel reads one. Each is documented where it bit, because the symptom is
always a plausible number rather than a failure.

The worst of them was in the comparison itself: the baseline's *bf16* weight
was never rotated, while `native_route_plan` picks `hipblaslt_bf16` -- which
reads exactly that tensor -- on most large-M points. Those rows measured a hot
baseline against a cold mori, which ran against mori, not for it. Corrected,
SGLang is 3.6-18.6% slower there and mori's margin moves up to 28 points in its
favour. The rows whose route never reads that tensor are unchanged within a
point, which is how the fix was confirmed to touch only what it should.

**Failure has to reach the caller.** A correctness sweep that printed `fail=N`
and exited 0 -- verified with a probe rigged to fail all 112 cases -- is now a
pytest grid. Every benchmark returns non-zero, `sweep.py` exits with the count
of failed points, and "mori declines this shape by construction" is recorded as
a result rather than a crash, so ten expected declines cannot bury one real
break.

**Six entry points, not seventeen.** `bench_gemm.py` covers the two scopes that
are genuinely different questions -- `--scope kernel` for pre-quantised operands
and `--scope linear` for a whole layer -- where three scripts had encoded the
distinction implicitly. Six shell sweeps and a tuning script become `sweep.py`
presets; two reporters become `report.py`, which keys on every axis a sweep
declares varying and so cannot silently overwrite one configuration with
another. The SGLang baseline is optional throughout: mori's own numbers need
nothing but mori.

One historical experiment is kept and renamed: `fused-blockscale-control` ran
V4.1-Flash's shapes under a quantisation the checkpoint does not use, while
presenting itself as a V4.1 benchmark.

`--mode split-lsa --gather-dtype fp8` is refused where it is asked for. The LSA
all-reduce has no fp8 leg, and the combination used to run a bf16 collective and
report it under the fp8 label; the validation gate's two-sided band caught it,
but only after paying for the run.

`docs/` keeps how-to-run and a summary; every full table lives in the operator
README, next to the code it is about.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
The SGLang-side numbers were taken against a synthetic snapshot of the
container image, because the V4.1 code existed on no branch at the time. It
does now -- sgl-project/sglang#39857, `kevin-mii:dsv41-amd-main` -- so all
three servers were re-run on top of it.

The fused wo_b reproduces and is slightly better: fp8 wire +5.6% to +6.0% at
bs 4-16 against +5.1% to +5.7% before, bf16 +1.9% to +2.8% against +2.3% to
+2.6%. GSM8K 0.917 / 0.912 / 0.918.

Added, because it had never been in this section: the standalone GEMM on its
own, +2.2% to +2.5% with nothing fused, GSM8K 0.920 / 0.928. And a note that
the two results do not add -- `fused_wo_b` runs first and the linear only sees
what it declines, so they split `wo_b` rather than stacking on it. Both
together is a fourth configuration nobody has measured.

**One claim here was wrong and is corrected.** It called bs=1 a built-in
control on the grounds that M=4096 is below both floors. The floors are 8192
for bf16 and 2048 for fp8, and the shape log confirms the consequence: the
bf16 wire's smallest served M is 7734, the fp8 wire's is 1792. bs=1 is a
control for the bf16 column only, and the fp8 column's +3.4% there is a real
win that was being dismissed as drift.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
Both SGLang-side paths enabled together, which this section had been calling
unmeasured: +7.4% to +7.7% prefill throughput at bs 4-16, against +6.0% for the
fp8 wire alone and +2.3% for the standalone GEMM alone. They stack without
adding, and the structure says why -- `fused_wo_b` runs first and the linear
only sees what it declines, so at `wo_b` they divide the layer by M, and the
~1.7% on top is the column-parallel layers fusing cannot reach.

Also recorded: one server out of nine hit `hipErrorIllegalAddress` under
sustained load, and re-running the same variant completed cleanly. Written as
what it is -- one sample, cause not established, the combination merely the
only configuration that has shown it -- rather than as a known defect or as
nothing. The five hypotheses already eliminated are listed so the next person
does not re-walk them, and so is the reason `AMD_SERIALIZE_KERNEL=3` is being
held back: it is worth spending on a repro that reproduces, and serialising
perturbs the timing a race would depend on.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
…lished

Rebasing onto main brought in #688, which converted this file's one link into
`python/` to an absolute GitHub URL: Sphinx builds `docs/` alone, so a relative
path to a file outside that tree is a broken reference and `-W` fails the build.

The rewrite in this branch added three more of them. Same rule applied to all
four. Links the other way -- from the operator README into `docs/` -- stay
relative, because that file is not in the toctree and GitHub resolves them.

Verified with the gate itself rather than by eye:
`sphinx -E -n -W --keep-going` builds clean, and `tools/check_docs_links.py`
reports 8872 local links across 25 pages with 0 errors.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
@yangyuhuiling
yangyuhuiling force-pushed the xiangch/gemm-ar-dsv41-shapes branch from ce1c5f7 to 0fda3c1 Compare September 20, 2026 08:58
Two unrelated CI failures on this branch. They are in one commit because
`black` reflowed lines adjacent to the real fix, so splitting them would mean
hunk surgery on the same few lines rather than a cleaner history.

**The CCO unit test.** All ten `test_standalone_mxfp8_gemv` cases failed with

    AttributeError: 'Int32' object has no attribute 'type'. Did you mean: 'dtype'?
    _arith_ops_gen.py:3163, in ShRUIOp.__init__

while the same tests passed here. The generated MLIR builders infer their
result type as `operands[0].type`, so they need an operand that is already an
`ArithValue`; `fx.Int32` is a wrapper and carries `dtype`. FlyDSL's operators
go through `_make_binop` -> `_extract_arith`, which unwraps first, so
`(v >> shift) & 0xFF` is right on any FlyDSL that has the operators at all.

Whether the wrapper exposes `.type` differs between the FlyDSL here and the
one in CI's container, which is why this was invisible locally. The rule that
holds on both: hand wrappers to FlyDSL operators, never to raw builders.

`kernels_fused._load_byte` carried the identical expression and the identical
latent break; CI never reached it because the GEMV tests failed first.

The one `arith.andi` over two predicates became nested `arith.select` -- same
result, and `select` is what the rest of mori builds with, so it is the op
proven against whatever FlyDSL CI ships. Every `arith.*` call left in this
branch is now one `origin/main` also uses.

146 tests pass locally, including the fp32-reference numerics that would catch
any change in what the shift and mask compute.

**pre-commit.** Formatting from the hooks themselves, plus two they flagged
and could not fix:

- `needs_sglang` in bench_gemm.py was dead (F841). It computed a condition
  nothing read; `build_mxfp8(want_sglang=...)` already gates the only thing
  that needs it, and `--scope linear` imports SGLang's quantiser inside
  `mori_linear_call`, where a missing build is a reported per-impl failure and
  a non-zero exit. Deleted rather than wired up: there was no guard to restore.
- `l` as a comprehension variable in sweep.py (E741) -> `line`.

The license hook appended the full MIT header to timing.py above its SPDX
short form, leaving two notices; dropped the short one. The second notice in
`_gemm_a8w8_8wave.py` is the deliberate vendored-from-aiter attribution and is
left alone.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>

@jhchouuu jhchouuu 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.

Reviewed on 8x MI355X against the PR head. Four things to fix, all on the new public surface or in the benchmark contracts — none of them in the MXFP8 kernel numerics, which held up everywhere I pushed them. Details inline.

What checked out

  • The GEMV masked-step path is correct well beyond the two shapes the grid test covers: 59 cases across NSTEPS 1–16, including waves > NSTEPS where 15 of 16 waves are entirely out of range. All within the fp8 floor, no NaN. Given the history in that code, this is the result that matters most.
  • _WIDE_TILE_MIN_GRID = 140 picks the faster tile on 16/16 measured points, and both tables in gemm.py reproduce within 1–2% (M=16384: 165.9/240.0 and 146.9/212.0 against your 167.2/239.1 and 148.6/212.0). The grid-based threshold is doing what the M-based one could not.
  • The GEMV tuned entries beat what _HEURISTIC would pick on every bucket, 6.7% to 28.8%.
  • wo_a: a config sweep finds 13.9–17.5% over the heuristic at M=8/32 — consistent with, and slightly better than, the 12–16% the docstring admits to. Two _tune() entries would recover it.
  • TP4 mxfp8, M=16384: fused+fp8 against split+bf16 is −24.0%, spread ≤0.15% over repeats. Fusion alone is 6.0% (bf16) and 9.3% (fp8); the rest is the wire, which also takes relL2 from 2.4e-3 to 2.3e-2. The PR is clear about that split, and it holds.

Not covered: anything comparing against SGLang (this environment is missing mori_mxfp8_common.py), and end-to-end serving — so the unreproduced hipErrorIllegalAddress in the docs is still open.

# _LaneTransposeStoreC's permlane mapping pairs exactly two N-tiles. So 128
# is legal without permlane and not with it, which is what the store's own
# assert says a few frames deeper and less usefully.
n_floor = 256 if permlane else 128

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.

The permlane guard is a lower bound where the store needs equality.

n_floor admits any BLOCK_N >= 256, but _LaneTransposeStoreC._emit asserts n_tiles_b == 2, i.e. BLOCK_N == 256 exactly. So block_n=512 passes here and dies a few frames deeper:

supports_gemm(n=8192, k=1280, block_n=512) -> True
  M=   64 narrow=True  rel_l2=1.659e-03  ok
  M= 1024 narrow=True  rel_l2=1.661e-03  ok
  M= 4096 narrow=False AssertionError: the permlane mapping pairs exactly two
                       N-tiles (BLOCK_N == 256); got n_tiles_b=4

An instance that serves small batches fine and fails once the batch grows is the worst shape for this to take.

Suggest BLOCK_N == 256 on the permlane branch rather than >= n_floor. Fixing it here covers GemmAllReduceOp too, rather than adding a third copy of the rule at each entry point.

k: int,
block_n: int = DEFAULT_BLOCK_N,
):
why = _tile_constraints(MXFP8_BLOCK_M, block_n)

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.

supports_gemm returns True for block_n=512 on shapes the wide path cannot actually store (see the note on kernels_fused.py:1140). If the guard there is tightened to == 256 this predicate follows automatically; if not, it needs its own check, since it is the documented way to ask.

(No issue with the narrow tile being checked here and not in __init__ — TILE_N_GRANULE=256 makes N % block_n == 0 imply N % 128 == 0, so that gap is unreachable.)


got = call(weights)
torch.cuda.synchronize()
if ref is None:

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.

The first implementation becomes its own reference, so it is never validated.

ref is seeded from the first got, making rel_l2 identically 0 for it and validated unconditionally true. The rel_l2 is None branch passes for the same reason. A MORI-only run — which the docs support — reports validated=true whatever it computed.

Later impls are still compared against the first, so a bug in one variant is caught; a bug shared by all of them, or a single-impl run, is not.

bench_gemm_ar.py:210 already has reference_partial() — an fp32 reference matching the same quantisation — used with a real tolerance at :486 and :506. Reusing it here seems better than writing a second one, and would make the new benchmark at least as strong as the existing one next to it. Failing that, validated should be None when there is nothing to compare against, not True.

)
failures += 1
continue
ok = rel is None or rel <= args.tol

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.

Same as bench_gemm.py:412: with --baseline none there is no reference, rel stays None, and that is treated as a pass.

row["route"] = native_route_plan(
m, n, k, ops["layer"].weight_bf16 is not None, False
)
call, weights = sglang_linear_call(ops["layer"], x)

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.

--scope kernel with --impl sglang compares two different inputs.

sglang is in the default --impl list and --scope kernel does not exclude it, but the two branches do not share operands: mori_kernel_call builds its own fp8 A and scales from seed 99 (:256) while the sglang branch runs sglang_linear_call(layer, x) on the outer bf16 x.

Since auto comes first and becomes ref, the effect is that sglang gets rel_l2'd against an unrelated result and reads as failing validation — a false accusation against the baseline rather than a false pass. The row also reports includes_quant=false (:381) while that branch does quantise, so the timing is not comparable either.

Either reject the combination, or have both sides consume the same pre-quantised input with quantisation outside the timed region.

Comment thread benchmark/cco/flydsl/gemm_ar/sweep.py Outdated
world=4,
grid=_grid(
m=[4096, 8192, 16384],
mode=["gemm-only", "split-sdma", "split-lsa", "fused-sdma", "fused-lsa"],

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.

split-lsa is in this preset's mode list and the preset appends --gather-dtype fp8 to every mode, but bench_gemm_ar.py:281 rejects exactly that pair before the communicator is up:

--mode split-lsa has no fp8 gather leg (build_lsa_ar takes no gather_dtype);
use split-sdma or fused-sdma for --gather-dtype fp8

sweep.py fused-fp8 --dry-run emits three such commands (m=4096/8192/16384), so 3 of the 15 points cannot pass even on a healthy box. Suggest dropping the mode from this preset, or skipping it explicitly as unsupported.

All four reproduce as reported; none are in the MXFP8 numerics.

**`supports_gemm` said yes to a tile the store cannot write.** The permlane
guard was `BLOCK_N >= 256`, but `_PermlaneStoreC._emit` pairs *exactly* two
N-tiles and asserts `n_tiles_b == 2`. So `block_n=512` passed every predicate
and died several frames deeper -- and because the wide tile is only chosen once
the grid is large enough, such an instance serves small batches correctly and
fails when one grows. That is the worst shape a failure can take.

`GemmAllReduceOp` already rejected it at its constructor; `Mxfp8GemmOp` and
`supports_gemm` are new here and did not. Rather than restate the rule a third
time, it is now `permlane_tile_constraint()` next to `_tile_constraints` in
op.py, used by all three, with the kernel's own assert tightened to an equality
as the backstop for callers that compile directly. Review suggested the assert
alone would cover the predicate; it does not -- `supports_gemm` has to answer
without compiling, which is why the rule has to exist as a predicate too.

**The benchmarks reported `validated=true` for runs that validated nothing.**
`bench_gemm.py` seeded `ref` from the first implementation's output, so that
impl was its own reference: `rel_l2` identically 0 and a pass regardless of
what it computed. A single-impl run -- which the docs recommend -- checked
nothing, and a bug shared by every impl was invisible. `bench_gemv.py` had the
same hole via `--baseline none`.

Both now score against `reference_partial()`, the fp32 reference
`bench_gemm_ar.py` next door already uses, built from the raw operands. The
first impl's relL2 is now 1.66e-03 rather than 0, which is the fp8 floor and
matches what review measured independently. Where no independent reference
exists -- anything quantising inside the timed region -- `validated` is `None`
and prints as `n/a`, never `True`.

**`--scope kernel --impl sglang` compared two different inputs.** mori's kernel
path builds its own pre-quantised A from seed 99 while SGLang's entry point is
a linear and quantises the outer bf16 `x`. The baseline was scored against an
unrelated result and read as failing validation -- a false accusation -- and
the timings were not comparable either, only one side carrying the
quantisation. Now refused at argument parsing, which is what this file's own
docstring already said: linear is the scope where that comparison means
anything.

**`sweep.py fused-fp8` shipped 3 points that cannot run.** The preset listed
`split-lsa` while appending `--gather-dtype fp8`, and the bench refuses that
pair at entry -- a rejection this branch added and did not propagate here.
15 points, 3 of them unrunnable on a healthy box. Dropped from this preset
only; the bf16 preset keeps `split-lsa`, where it is legal.

Not addressed: the `wo_a` tuning entries review found worth 13.9-17.5%. That is
a new performance claim and wants its own reproduction rather than being taken
on the review's numbers inside an already-large PR.

66 tests pass; both kernel-scope paths verified against the new reference.

Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
@yangyuhuiling yangyuhuiling changed the title gemm_ar: MXFP8 (32-wide ue8m0) support, and the GEMM usable without a collective feat(cco): gemm_ar MXFP8 support, and the GEMM usable without a collective Sep 21, 2026

@jhchouuu jhchouuu 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.

LGTM

great work!

@yangyuhuiling
yangyuhuiling merged commit bb3a9aa into main Sep 21, 2026
15 checks passed
@yangyuhuiling
yangyuhuiling deleted the xiangch/gemm-ar-dsv41-shapes branch September 21, 2026 10:40
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants