Skip to content

fix(cuda/moe): tag the pre-Ampere grouped GEMM arm Sm70, not Sm75 - #1562

Merged
inureyes merged 2 commits into
mainfrom
fix/issue-1544-grouped-gemm-arch-tag
Sep 1, 2026
Merged

inureyes merged 2 commits into
mainfrom
fix/issue-1544-grouped-gemm-arch-tag

Conversation

@inureyes

@inureyes inureyes commented Aug 31, 2026 •

Copy link
Copy Markdown
Member

Summary

dispatch_cutlass_arch in grouped_gemm_unaligned.cu mapped every device below compute capability 8.0 to cutlass::arch::Sm75, which names the m16n8k8 MMA Turing introduced and an sm_70 part does not have. The pre-Ampere arm now selects cutlass::arch::Sm70, and the Sm75 placeholder that initialised fun in get_grouped_mm_funcion is gone.

#1544 was written verification-first and said a negative result would be a complete answer. The answer is not negative, and not the positive one either.

The branch is live, on the checkpoint this epic benchmarks

The issue expected grouped_gemm_v2 to be unreachable because mlxcel's fused MoE path bypasses it. That is right about decode and wrong about prefill. #629's sorted-MoE prefill fast path (quantized.cpp:359) routes a quantized GatherQMM into cutlass_grouped_gemm_unaligned once B >= min_rows * num_experts; for gemma-4-26b-a4b-it-4bit that is E = 128 and top_k = 8, so the gate is a 128-token prompt. Three nsys runs, same binary, same checkpoint:

Run Prompt B cutlass::Kernel<GemmGrouped> prepare_grouped_mm_data
default 573 tokens 4584 180 inst, 853.8 ms, 3.8% of GPU time 180 inst
MLXCEL_GATHER_QMM_GROUPED=0 573 tokens 4584 absent absent
default 46 tokens 368 absent absent

The 46-token run reproduces the #1538 baseline's MoE profile exactly at 326 qmm_naive instances, which is the cross-check that it is the same measurement the baseline made. That baseline profiled MoE below the gate, which is why no CUTLASS kernel appears in it; its MoE section is annotated here so the omission does not keep reading as absence.

Also reachable, without any gate, by every non-quantized MoE checkpoint through SwitchLinear::Regular::forward (gpt_oss, deepseek, qwen3_next, jamba, nemotron_h, exaone_moe and the rest).

Nothing was computing wrong, and that was measured

The pre-Ampere arm resolves to GemmConfiguration's primary template, which is OpClassSimt with InstructionShape<1, 1, 1>. Both tensor-core specializations are constrained on Arch::kMinComputeCapability >= 80, so no pre-Ampere tag can select an MMA atom of any shape, with or without MLX_ENABLE_TF32. CUTLASS therefore erases the tag, and the kernel that actually ran says so in its own name: MmaPipelined / MmaSimt / OpMultiplyAdd, with no Sm70, Sm75 or Sm80 token anywhere in it.

Correctness was checked rather than inferred, because a mismatched CUTLASS path can return plausible numbers that are wrong and greedy decoding will not reveal it. Two independent checks: the grouped path and the legacy qmm_naive reference path produce a byte-identical 1,793-character greedy continuation on a 285-token prompt, and this holds for all four combinations of {before, after} x {grouped, legacy}; and new unit tests compare gather_mm against an f64 dense per-expert reference on the model's real expert dims. The pre-existing test_gather_mm asserted the output shape and never looked at a value, which is how this went unchecked.

The retag moves no device code, on any architecture

Compiling the translation unit before and after with the production flags, and comparing cuobjdump --dump-sass per symbol:

Target Symbols before Symbols after Only in one Bodies differing SASS compared
compute_70 51 51 0 0 58,211,476 bytes
compute_80 51 51 0 0 35,077,812 bytes
compute_121 51 51 0 0 50,952,400 bytes

This is #1539's technique, and unlike #1541's case it transfers: this file compiles its CUTLASS kernels ahead of time, so its object holds real device code and an identical dump is a result rather than a tautology. Whole-file dumps differ only in cubin emission order, so the comparison is per symbol.

A control settled the design against giving Turing its own arm: adding an Sm70 arm alongside Sm75 emits the identical 51 device symbols and a byte-identical 493,538-line dump, for 26 extra host-side instantiations and 194,704 more bytes of object. The pre-Ampere arm therefore stays one arm, and its tag names the floor of the range it covers.

What changed

  • patches/mlx/backend/cuda/gemms/grouped_gemm_arch.h (new): the architecture decision as a pure constexpr function of the compute capability major version, with the evidence for why one arm covers Volta and Turing and what has to change before that stops holding.
  • patches/mlx/backend/cuda/gemms/grouped_gemm_unaligned.cu: dispatch_cutlass_arch switches on that function; get_grouped_mm_funcion returns a named GroupedGemmFn starting from nullptr with a loud guard instead of an Sm75 placeholder instantiation; two static_asserts pin the pre-Ampere configuration as SIMT (including under kEnableTF32) and its stage count at 2.
  • src/lib/mlxcel-core/cpp/grouped_gemm_arch_probe.cpp (new) and build.rs: C shim over the shipped decision, compiled unconditionally rather than behind cuda, following the perf(cuda/quant): size the qmm_naive tile from the device shared-memory budget, not an sm80 flag #1541 precedent.
  • src/lib/mlxcel-core/src/grouped_gemm_arch_tests.rs (new): enumerates the tag mapping over every compute capability, on the shipped function, with no GPU.
  • src/lib/mlxcel-core/src/grouped_gemm_numeric_tests.rs (new): gather_mm against an f64 dense per-expert reference at k = 2816, n = 704 and its transpose, across both entry points, both kAlignmentC arms, both operand layouts and f32/bf16/f16, plus a constant-per-expert case that catches a mis-gathered index.
  • docs/benchmark_results/grouped-gemm-arch-v100-2026-08-31.md (new), the test(benchmark): Volta (sm_70) baseline record and build coverage #1538 baseline's reserved post-program row and its MoE section, and a CHANGELOG entry.
  • TECHNICAL_REPORTS/1562-grouped-gemm-arch-tag-sm70-20260901.{en,ko}.md (new), the bilingual report.

kStages and cp.async

The issue asked whether kStages = 3; // use SM80_CP_ASYNC can reach a pre-Ampere arch tag. It cannot, structurally: that member belongs to GemmConfiguration<float, cutlass::arch::Sm80, kAlignmentC, true>, an explicit full specialization on Sm80 that no other tag can name. The pre-Ampere arm gets kStages = 2 and runs MmaPipelined, which the profiled kernel name confirms. Now asserted at compile time rather than left as reasoning.

The cuFuncSetAttribute gap handed over from #1541: recorded, not fixed

#1541 handed this issue the missing dynamic shared-memory opt-in in gather_gemm.cu. Two findings, and the conclusion is to record.

The gap is in the shared encoder rather than one launch site: CommandEncoder::add_kernel_node_raw sets sharedMemBytes and calls cudaGraphAddKernelNode without ever calling cudaFuncSetAttribute, so every launch site must opt in itself. But the reachable Volta configuration is nowhere near the ceiling, measured from the profile rather than computed: the grouped GEMM asks for 10,320 bytes of dynamic shared memory against the 49,152-byte non-opt-in limit, 21% of it. gather_gemm.cu also remains unreachable, which reproduces here (nm -C finds zero undefined references to mlx::core::gather_mm(bool, bool, ...) in the built libmlx.a, while all three grouped_gemm_unaligned.cu entry points carry one from matmul.cpp.o).

What is genuinely open and new: the sm_80 tensor-core configurations in this same file use much larger tiles and, in the tf32 arm, three stages. That is the plausible place for this path to trip the ceiling, and it needs an Ampere-or-later part to answer. Left as a follow-up rather than fixed blind.

MoE decode and prefill on Volta, before and after

Five repetitions each, warm PTX cache, nvidia-smi --query-compute-apps asserted empty before every run, decode as a slope over -n 40 to -n 120 per the baseline's rule 1, prefill at -n 1 on the 285-token prompt that crosses the grouped-GEMM gate.

Checkpoint Metric Before After Delta Before-arm spread
gemma-4-26b-a4b-it-4bit decode 29.66 ms/tok (33.72 tok/s) 29.73 ms/tok (33.63 tok/s) +0.26% 10.42%
gemma-4-26b-a4b-it-8bit decode 31.68 ms/tok (31.56 tok/s) 32.71 ms/tok (30.57 tok/s) +3.24% 11.11%
gemma-4-26b-a4b-it-4bit prefill 12,720 ms (22.41 tok/s) 12,644 ms (22.54 tok/s) -0.59% 0.75%
gemma-4-26b-a4b-it-8bit prefill 15,212 ms (18.73 tok/s) 15,606 ms (18.26 tok/s) +2.59% 14.15%

Every delta is smaller than the before arm's own repetition spread, which is the only outcome byte-identical device code permits. The 8-bit prefill before cell's 14.15% spread is one cold first repetition that is left in rather than dropped; comparing repetitions 2 to 5 on both sides gives 15,633 ms against 15,606 ms, a -0.17% delta.

The before column is re-measured on this host from this worktree rather than quoted from #1538, whose 38.61 and 34.29 ms/token predate #1539. It also inverts that record's 4-bit against 8-bit MoE finding, in the same direction #1539's own record reports for the dense pair.

Deferred to GB10

No sm_80-or-later device exists on this host, so the following are left unticked on #1544 rather than claimed. See epic #1536's ## GB10 (sm_121) continuation.

  • GB10 MoE output byte-identical: not run. What is recovered locally is the mechanism rather than the measurement: the tag mapping is provably unchanged for every compute capability at or above 8, and the device code that mapping selects is byte-identical before and after at compute_80 and compute_121, per symbol. That leaves no route by which GB10 output could move, but identity is a measurement and this is an argument.
  • GB10 MoE throughput unmoved: same, and for the same reason.
  • cargo test --features cuda green on sm_121: needs a GB10 host. On sm_70 the mlxcel-core library suite is 1672 passed against 1 pre-existing failure that this change is shown not to cause; see the test plan.
  • The sm_80 tensor-core grouped-GEMM configurations against the 48 KB dynamic shared-memory ceiling: newly raised here, genuinely open, and worth its own issue.

Note that the CUDA sm_70 compile CI check is not evidence for any of this. CUDA 13 removed Volta support and cannot compile compute_70, so that job passes in about 11 seconds by skipping; the local build is the only real coverage.

Test plan

  • MLX_CUDA_ARCHITECTURES=70 cargo build --release --features cuda, both arms; cuobjdump --list-elf libmlx.a reports 96 cubins, every one sm_70.
  • cargo test -p mlxcel-core --release --features cuda --lib -- --test-threads=1 on the V100: 1672 passed, 1 failed, 1 ignored in 496 s. The one failure is sampling::tests::temperature_one_support_unchanged, and it is pre-existing on c2e54939 rather than caused by this change. That was isolated rather than assumed: reverting grouped_gemm_unaligned.cu to its base version, rebuilding, and re-running the test reproduces a byte-identical failure. It asserts bit-exactness of fused_sample_probs at T = 1.0 and fails by 1 ULP on 5 of 64 entries, the same class of sm_70 float-reduction non-determinism perf(cuda/quant): accumulate qmv in float below Ampere for bf16 #1557 recorded for tests/cuda_qmm_determinism.rs. This branch touches no sampling code. Worth its own issue; flagged here rather than folded into this one. (--test-threads=1 is what fix(test): the full mlxcel-core lib suite aborts under CUDA graph capture, so the real merge gate does not run #1048 requires of the CUDA suite.)
  • cargo clippy -p mlxcel-core --release --features cuda --lib --tests -- -D warnings clean.
  • cargo fmt -p mlxcel-core -- --check clean.
  • Reachability: three nsys profiles with --cuda-graph-trace=node, table above.
  • Correctness: byte-identical greedy continuation, grouped against legacy, before against after, all four combinations.
  • Device-code delta: per-symbol SASS comparison at compute_70, compute_80 and compute_121.
  • Throughput: five repetitions per cell, both MoE arms, decode slope and prefill, table above.

Closes #1544

`dispatch_cutlass_arch` in `grouped_gemm_unaligned.cu` mapped every device below compute capability 8.0 to `cutlass::arch::Sm75`, which names the `m16n8k8` MMA Turing introduced and an sm_70 part does not have, and `get_grouped_mm_funcion` opened with a matching `Sm75` placeholder. The pre-Ampere arm now selects `cutlass::arch::Sm70`.

The branch is live, and on the checkpoint epic #1536 benchmarks. #629's sorted-MoE prefill fast path routes a quantized `GatherQMM` into `cutlass_grouped_gemm_unaligned` once `B >= 8 * num_experts`, which for `gemma-4-26b-a4b-it-4bit` means a prompt of 128 tokens or more. An nsys profile at 573 prompt tokens shows 180 `cutlass::Kernel<GemmGrouped>` launches taking 3.8% of GPU time; the same profile with `MLXCEL_GATHER_QMM_GROUPED=0` shows none, and so does the 46-token prompt the #1538 baseline profiled, which is why nobody had seen it. That baseline's MoE section is annotated with the gate so the omission does not read as absence.

Nothing was computing wrong. The pre-Ampere arm resolves to `GemmConfiguration`'s primary template, which is `OpClassSimt` with `InstructionShape<1, 1, 1>`, so no arch tag can reach an MMA atom there and CUTLASS erases it: the kernel that runs is an `MmaSimt` / `OpMultiplyAdd` / `MmaPipelined` instantiation carrying no architecture token in its name at all.

Correctness was measured rather than assumed, because a mismatched CUTLASS path can return plausible wrong numbers that greedy decoding hides. The grouped path and the legacy `qmm_naive` path produce a byte-identical 64-token greedy continuation, before and after this change alike. New tests compare `gather_mm` against an f64 dense per-expert reference on the model's real expert dims across both entry points, both output-alignment arms, both operand layouts and f32, bf16 and f16; the pre-existing `test_gather_mm` asserted only the output shape.

The retag moves no device code. Compiling the translation unit before and after at `compute_70`, `compute_80` and `compute_121` yields the same 51 device symbols with byte-identical SASS bodies at all three, 144 MB of dump compared per symbol. That is also what settled the design against a separate Turing arm: adding one emits the same device code twice for 26 extra host instantiations and 194,704 bytes of object.

The decision now lives in `gemms/grouped_gemm_arch.h` as a pure function of the compute capability major version, enumerated over every architecture by `grouped_gemm_arch_tests.rs` through a C shim with no GPU involved, which closes this issue's "zero change on sm_80+" criterion locally instead of deferring it to GB10. Two `static_assert`s pin the preconditions: that the pre-Ampere configuration is still SIMT, without which one arm could not cover Volta and Turing together, and that its stage count is 2, since the 3-stage `SM80_CP_ASYNC` pipeline is bound to an explicit `cutlass::arch::Sm80` specialization and `cp.async` does not exist before Ampere.

MoE decode and prefill are unmoved on a V100, as byte-identical device code requires: 4-bit decode 29.66 to 29.73 ms/token, 8-bit 31.68 to 32.71, 4-bit prefill 12,720 to 12,644 ms, every delta inside the before arm's own repetition spread. Full record in `docs/benchmark_results/grouped-gemm-arch-v100-2026-08-31.md`.

Refs #1536

Closes #1544
@inureyes inureyes added type:bug Bug fixes, error corrections, or issue resolutions priority:medium Medium priority area:core mlxcel-core: MLX FFI, primitives, KV cache, layers arch:moe Sparse mixture-of-experts decoder platform:linux Linux (CUDA / packaging) specific status:review Under review labels Aug 31, 2026
@inureyes inureyes added status:done Completed and removed status:review Under review labels Sep 1, 2026
@inureyes
inureyes merged commit e5cae85 into main Sep 1, 2026
13 checks passed
@inureyes
inureyes deleted the fix/issue-1544-grouped-gemm-arch-tag branch October 6, 2026 12:12
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

arch:moe Sparse mixture-of-experts decoder area:core mlxcel-core: MLX FFI, primitives, KV cache, layers platform:linux Linux (CUDA / packaging) specific priority:medium Medium priority status:done Completed type:bug Bug fixes, error corrections, or issue resolutions

Projects

None yet

Development

Successfully merging this pull request may close these issues.

fix(cuda/moe): grouped GEMM selects cutlass::arch::Sm75 on an sm_70 part

1 participant