Repository navigation
fix(cuda/moe): tag the pre-Ampere grouped GEMM arm Sm70, not Sm75 - #1562
Merged
Merged
Conversation
`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
7 of 13 tasks
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
dispatch_cutlass_archingrouped_gemm_unaligned.cumapped every device below compute capability 8.0 tocutlass::arch::Sm75, which names them16n8k8MMA Turing introduced and an sm_70 part does not have. The pre-Ampere arm now selectscutlass::arch::Sm70, and theSm75placeholder that initialisedfuninget_grouped_mm_funcionis 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_v2to 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 quantizedGatherQMMintocutlass_grouped_gemm_unalignedonceB >= min_rows * num_experts; forgemma-4-26b-a4b-it-4bitthat isE = 128andtop_k = 8, so the gate is a 128-token prompt. Three nsys runs, same binary, same checkpoint:Bcutlass::Kernel<GemmGrouped>prepare_grouped_mm_dataMLXCEL_GATHER_QMM_GROUPED=0The 46-token run reproduces the #1538 baseline's MoE profile exactly at 326
qmm_naiveinstances, 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_moeand the rest).Nothing was computing wrong, and that was measured
The pre-Ampere arm resolves to
GemmConfiguration's primary template, which isOpClassSimtwithInstructionShape<1, 1, 1>. Both tensor-core specializations are constrained onArch::kMinComputeCapability >= 80, so no pre-Ampere tag can select an MMA atom of any shape, with or withoutMLX_ENABLE_TF32. CUTLASS therefore erases the tag, and the kernel that actually ran says so in its own name:MmaPipelined/MmaSimt/OpMultiplyAdd, with noSm70,Sm75orSm80token 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_naivereference 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 comparegather_mmagainst anf64dense per-expert reference on the model's real expert dims. The pre-existingtest_gather_mmasserted 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-sassper symbol:compute_70compute_80compute_121This 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
Sm70arm alongsideSm75emits 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 pureconstexprfunction 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_archswitches on that function;get_grouped_mm_funcionreturns a namedGroupedGemmFnstarting fromnullptrwith a loud guard instead of anSm75placeholder instantiation; twostatic_asserts pin the pre-Ampere configuration as SIMT (including underkEnableTF32) and its stage count at 2.src/lib/mlxcel-core/cpp/grouped_gemm_arch_probe.cpp(new) andbuild.rs: C shim over the shipped decision, compiled unconditionally rather than behindcuda, 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_mmagainst anf64dense per-expert reference atk = 2816, n = 704and its transpose, across both entry points, bothkAlignmentCarms, 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.kStagesandcp.asyncThe issue asked whether
kStages = 3; // use SM80_CP_ASYNCcan reach a pre-Ampere arch tag. It cannot, structurally: that member belongs toGemmConfiguration<float, cutlass::arch::Sm80, kAlignmentC, true>, an explicit full specialization onSm80that no other tag can name. The pre-Ampere arm getskStages = 2and runsMmaPipelined, which the profiled kernel name confirms. Now asserted at compile time rather than left as reasoning.The
cuFuncSetAttributegap 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_rawsetssharedMemBytesand callscudaGraphAddKernelNodewithout ever callingcudaFuncSetAttribute, 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.cualso remains unreachable, which reproduces here (nm -Cfinds zero undefined references tomlx::core::gather_mm(bool, bool, ...)in the builtlibmlx.a, while all threegrouped_gemm_unaligned.cuentry points carry one frommatmul.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-appsasserted empty before every run, decode as a slope over-n 40to-n 120per the baseline's rule 1, prefill at-n 1on the 285-token prompt that crosses the grouped-GEMM gate.gemma-4-26b-a4b-it-4bitgemma-4-26b-a4b-it-8bitgemma-4-26b-a4b-it-4bitgemma-4-26b-a4b-it-8bitEvery 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.compute_80andcompute_121, per symbol. That leaves no route by which GB10 output could move, but identity is a measurement and this is an argument.cargo test --features cudagreen on sm_121: needs a GB10 host. On sm_70 themlxcel-corelibrary suite is 1672 passed against 1 pre-existing failure that this change is shown not to cause; see the test plan.Note that the
CUDA sm_70 compileCI check is not evidence for any of this. CUDA 13 removed Volta support and cannot compilecompute_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.areports 96 cubins, every one sm_70.cargo test -p mlxcel-core --release --features cuda --lib -- --test-threads=1on the V100: 1672 passed, 1 failed, 1 ignored in 496 s. The one failure issampling::tests::temperature_one_support_unchanged, and it is pre-existing onc2e54939rather than caused by this change. That was isolated rather than assumed: revertinggrouped_gemm_unaligned.cuto its base version, rebuilding, and re-running the test reproduces a byte-identical failure. It asserts bit-exactness offused_sample_probsatT = 1.0and 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 fortests/cuda_qmm_determinism.rs. This branch touches no sampling code. Worth its own issue; flagged here rather than folded into this one. (--test-threads=1is 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 warningsclean.cargo fmt -p mlxcel-core -- --checkclean.--cuda-graph-trace=node, table above.compute_70,compute_80andcompute_121.Closes #1544