Repository navigation
feat(cco): gemm_ar MXFP8 support, and the GEMM usable without a collective - #690
Conversation
1a2d8fd to
5e6ca07
Compare
…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>
ce1c5f7 to
0fda3c1
Compare
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
left a comment
There was a problem hiding this comment.
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 > NSTEPSwhere 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 = 140picks the faster tile on 16/16 measured points, and both tables ingemm.pyreproduce 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
_HEURISTICwould 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 |
There was a problem hiding this comment.
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) |
There was a problem hiding this comment.
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: |
There was a problem hiding this comment.
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 |
There was a problem hiding this comment.
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) |
There was a problem hiding this comment.
--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.
| world=4, | ||
| grid=_grid( | ||
| m=[4096, 8192, 16384], | ||
| mode=["gemm-only", "split-sdma", "split-lsa", "fused-sdma", "fused-lsa"], |
There was a problem hiding this comment.
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>
What this adds
mori.ops.gemm_argains support for 32-wide ue8m0 (MXFP8), the quantisationDeepSeek-V4.1-Flash actually ships (
weight_block_size [32, 32],scale_fmt "ue8m0"), and two ways to use the GEMM without a collective attached.GemmAllReduceOpRowParallelLinearat prefill MMxfp8GemmOp(gemm.py)Mxfp8GemvOp(gemv.py)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
ColumnParallelLinearhas no all-reduce to fuse with at all.CDNA4's
v_mfma_scale_f32_16x16x128_f8f6f4takes the ue8m0 scales asinstruction operands, so the 32-wide form needs no dequantisation arithmetic
at all -- unlike the existing
_BlockScaleKpromote/rescale chain, which isthere 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:
Fused GEMM + all-reduce at V4.1-Flash's
wo_b, TP4, M=16384: 1192us against1592 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:
accesses row-major, +50-67% on the GEMM.
opsel_bon the scaled MFMA.-29.1%
SQ_INSTS_VMEMwith VALU flat, worth -5 to -7.5%.Correctness
tests/python/cco/test_gemm_ar_op.py, plus theexisting 2-rank worker cases.
tl.dot_scaled, mori agrees with it bit forbit -- same instruction, same operands, and ue8m0 scales are exact powers of
two. That is asserted, not just tolerated.
mxfp8_gemvat every M and both shapes.check_gemv_grid.sh) passes on both realshapes and every M in {1, 2, 8, 16, 32}.
Notes for review
_gemm_a8w8_8wave.pyis vendored from aiter and now diverges in two places,not one.
Mfma16x16x128grew the scaled form of the atom:opsel_b_per_tileon the constructor,
scale_a/scale_boncall. Both default to off and theunscaled path is byte-for-byte what it was, so a re-vendor stays a merge. The
file's header says so.
_PinnedLaunchreaches 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 tupleelement-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_unpackedand
mxfp8_roware the scale layoutsmxfp8was chosen over. They are ~5 linesof plumbing and they keep the choice reproducible; happy to drop them if you
would rather not carry them.
bench_gemm_ar.pykeeps its single-call graph timer deliberately. On thisbox 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.pyinstead, whichamortises 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.mdis now how-to-run plus asummary (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 raisesNotImplementedError. Its regions aresized but nothing writes the payload; that leg is the largest remaining lever
on the fused path.
gate_up(N=1152) is not expressible -- N is 4.5 tiles of256. Measured on the narrow tile anyway and it never wins, so
supports_gemmis left refusing it.