Repository navigation
[Metal] global scale for qmm - #4458
Merged
Merged
Conversation
nastya236
force-pushed
the
gather-qmm-nvfp4-global-scale
branch
from
September 3, 2026 23:45
8a420c9 to
f7c3404
Compare
zcbenz
approved these changes
Sep 7, 2026
zcbenz
left a comment
Member
There was a problem hiding this comment.
Nice work! just some nitpickings.
| bn, | ||
| wm, | ||
| wn, | ||
| true) |
Member
There was a problem hiding this comment.
I think you can simplify the big if clause with:
kernel = get_qmm_nax_kernel_wrapped(
...
global_scale.has_value());
Collaborator
Author
There was a problem hiding this comment.
You right! I just added global scale always false to affine quantization.
| ).reshape((E, 1, 1)) | ||
|
|
||
| for dtype in [mx.float32, mx.float16, mx.bfloat16]: | ||
| for transpose in [True, False]: |
Member
There was a problem hiding this comment.
You can reduce the indentations by using product, for example:
mlx/python/tests/test_quantized.py
Lines 230 to 235 in ce916db
inureyes
added a commit
to lablup/mlxcel
that referenced
this pull request
Sep 10, 2026
) ## Why The fp8 round-trip bound failed on every Metal host, M1 Ultra byte-identically to M5 Max, because the pinned MLX `9a795735` predates ml-explore/mlx#4353: Metal and CPU encoded the mxfp8 E8M0 block scale as `round(log2(amax / 448))`, so about half the blocks saturated their maxima, losing up to `1 - 2^-1/2`. CUDA rounds up, which is why #1742 passed on GB10. The test was right, and `requantize_block_fp8_weights`, the only E8M0 quantize caller, was clipping vendor FP8 checkpoints on Metal. Widening the bound, as the issue proposed, was rejected. ## What changed - MLX pin `9a795735` to upstream main `81ba1c6a` (99 commits). The seven overlays whose targets upstream touched are three-way merged and keep their deltas; the other 21 are unchanged upstream. `metal/compiled.cpp` also drops a leftover `elem_to_loc_1<uint>` that undid part of ml-explore/mlx#3720, so its only delta is the mixed-dtype cast. - Adaptations to what the new pin changes under the bridge: - `gather_qmm` gained `global_scale` ahead of `sorted_indices` (ml-explore/mlx#4458), so all 13 calls pass `std::nullopt`. - CUDA now needs cuSOLVER (ml-explore/mlx#4208). `link_cuda()` names it, `docs/installation.md` lists it, and CI's link job now runs on pin and CUDA link-list changes, which is how this slipped past CI. - `array::is_available()` now throws and clears a failed launch's error (ml-explore/mlx#3742). The rejection sampler's drain now reads status, signal and error pointer instead, so a GPU fault no longer terminates the process from inside `fused_sample` or hides the error from its owner. A regression test aborts the process with the old drain and passes with the new one. - The round-trip check moves into its own test, `fp8_block_requantize_round_trip_stays_within_half_an_e4m3_step`, with the same seed and shape. It states the derivation and asserts `group_max <= 448 * scale` per block. ## Validation (M1 Ultra, macOS 27.0) - fp8: the old pin fails the new test at block 0 (the maximum 4.8046875 scales to 615 and saturates). The new pin saturates 0 of 650 blocks, with a worst error of 0.0489 of the group max against 0.2928 before. - Turbo launchers pass (max RMS 1.7263e-4 and 1.5259e-4). - Teacher-forced logit traces, old pin vs new, on five checkpoints at widths 1, 8 and 256: 0 disagreements on decided positions. Over 4,096 positions on qwen3-30b-a3b, 1 of 1,720 decided positions differs, at the reference's rank 2, with perplexity -0.20%. - Branches the short runs missed: - head-dim-512 decode past 1,024 keys: identical text. - GQA-8 decode past 8,192 keys: identical text, decode 55.6 to 60.8 tok/s. - head-dim-72 vision towers: image prefill 239 to 279 tok/s. Descriptions diverge into equally faithful text; decided answers are unchanged. - Short-context throughput is within 0.6% on three checkpoints. - Workspace gate: 10986 passed, 0 failed, 359 ignored. Clippy, fmt, both pin parsers and the cross-repo reference check are clean. - CUDA was compiled and linked in CI but not run: `OpenXLA feature link` linked a `--features cuda,xla-iree` release binary on GB10 at the new pin, which covers the overlays and the cuSOLVER link. The green `CUDA sm_70 compile` check is a skip, because CUDA 13.0 cannot target sm_70. The per-checkpoint numbers, method and derivations are in `TECHNICAL_REPORTS/1772-mlx-pin-mxfp8-round-up-20260911.en.md`. ## Not validated, or reported and not fixed - No CUDA execution. No M5 Max run: the generation-17 NAX paths changed upstream in this range are unmeasured. - `array_evaluated_bytes`, the server's lookahead read, is another non-`Result` bridge function that now throws on a failed launch. It needs routing through the scheduler's step-failure path. - The mixed-dtype cast in `metal/compiled.cpp` would cast a comparison's inputs to `bool`. This is latent, since no compiled function contains a comparison. Closes #1769
aaishwarymishra
pushed a commit
to aaishwarymishra/mlx
that referenced
this pull request
Sep 11, 2026
2 tasks done
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.
This pull request adds global scales for MoE weights that are quantized to nvfp4 only for Metal.
Two important statements:
gather_qmmis used primarily on Metal for quantized MoEs.fp8_block_scalefor 16 elements andfp32_global_scaleper tensor scale.Then to dequantize the weight W fo 16 elements block w:
w = e2m1_code × fp8_block_scale × fp32_global_scaleNaive implementation would be to multiply the result of the multiplication:
x = x @ dequant(W) * fp32_global_scale
which is incorrect, because fp32_global_scale should be used in dequantization kernel to get a correct result
Note: it looks like a big diff but it is just adding new parameter to all related operations + changing dequantization part in
gather_qmm.Example:
TODO: add the same for CUDA for gather_qmv.