Repository navigation
perf(cuda/quant): block output rows per warp in qmv below Ampere - #1681
Merged
Merged
Conversation
`qmv` gives each warp one output row and has that warp read the whole activation vector, so a launch moves `n * k * 2` bytes of activations against `n * k / 2` bytes of 4-bit weights: four bytes of activation for every byte of weight. A controlled pair on a V100 puts a number on that. gemma-4-12B-it at 4 and 8 bits, same prompt, same 150 tokens, identical 49021 `qmv` launches, times taken as the union of `qmv` intervals inside the decode window: 2553.1 ms at 4 bits and 2959.2 ms at 8, or 52.1 us and 60.4 us per launch. The weight bytes go up 1.887x and the time goes up 1.159x. Both arms run the same number of dequantize-and-accumulate steps per row, because `elems_per_thread` is 16 at 4 bits and 8 at 8 bits, so a warp step moves 256 bytes either way. Solving `F + B = 52.1` and `F + 2B = 60.4` puts 43.8 us of a 4-bit launch, 84% of it, in work that does not scale with weight bytes. Decode below Ampere is not weight-bandwidth bound. So give each warp `R` consecutive output rows, load the activation tile once per k-step and apply it to all `R` weight rows. Weight traffic is unchanged and activation traffic divides by `R`. The arithmetic per row is untouched (same converter, same scale and bias order, same `cutlass::fma`), so the result is bit-identical to the stock kernel. The cost is on the other side: the same output rows are covered by `R` times fewer warps, and the accumulator array grows to `R * elems_per_thread`. Measured with `cuobjdump -res-usage` on the arm this card runs, 4-bit weights with `T = half_t`, nothing spilled: 40 to 48 registers at R = 1, 64 to 72 at R = 2, 96 to 142 at R = 4, which is 4 to 5, 3, and 1 to 2 resident blocks per SM. Decode as the slope between `-n 60` and `-n 200`, so the fixed first-token cost drops out: | | R = 1 | R = 2 | R = 4 | |---|---|---|---| | gemma-4-12B-it-4bit | 44.7 tok/s | 57.4 tok/s | 44.7 tok/s | | qwen3.8-27B-4bit | 20.6 tok/s | 24.9 tok/s | 24.7 tok/s | Generated text is byte-identical across all three at `--temp 0` on both checkpoints, which is the check that the tiling did not quietly change the arithmetic. R = 4 lands exactly back on R = 1 for the 12B and costs nothing on the 27B, and the register table says why: the 27B's projections are wider, so it still has enough blocks to fill the machine after the row count divides them. That is a shape dependence, not a rule, which is why R is a knob. The default is R = 2 and only on sm_70, read from the compute capability the encoder already carries. Turing stays at 1: it has a different register file per SM, no Turing device was available to measure, and the R = 4 reversal is a direct demonstration that this tradeoff does not extrapolate. `MLXCEL_QMV_ROWS_PER_WARP` overrides in both directions, and 1 routes back to the stock kernel with the stock grid. Also dispatches `max_x_rows` on the actual row count in the #725 multirow kernel, which had it pinned at 8 so a 2-row or 4-row launch still allocated eight rows of accumulators (95 to 109 registers against 55 to 56 and 68 to 71). That is worth 10.08 s to 9.20 s on a block-8 speculative run and nothing measurable at block 4, so it is included as a register-pressure cleanup rather than as the fix it first looked like. Refs #1543
Adds the row-blocking knob to the environment variable reference with the V100 numbers behind its default, and corrects the `MLXCEL_QMV_MULTIROW_MAX_ROWS` entry, which said the multirow kernel sizes its accumulators from a compile-time 8-row maximum. That is no longer true: the width is dispatched on the actual row count, so narrowing the window now actually returns registers. Refs #1543
The row-blocked `qmv` from the previous commit is bound by occupancy, not by the row count, and the register budget is the thing to set. Adding R = 3 and R = 6 to the dispatch turns the earlier two points into a curve, and the curve tracks the register column rather than the transaction count. Measured on a V100 at 4 bits on gemma-4-12B-it, decode as the slope between `-n 60` and `-n 200`: | R | tok/s | REG | blocks/SM | |---|---|---|---| | 1 | 44.7 | 40-48 | 4-5 | | 2 | 56.0 | 64-72 | 3-4 | | 3 | 53.4 | 80-110 | 2-3 | | 4 | 45.2 | 96-142 | 1-2 | | 6 | 42.7 | 128 | 2 | Every step up in R cuts transactions per output row, from 12 at R = 1 to 8, 6.67 and 6, and past R = 2 every step also costs a resident block per SM. The occupancy wins each time. So hand ptxas the budget directly. `__launch_bounds__(256, min_blocks)` caps registers at 65536 / (min_blocks * 256), and `MLXCEL_QMV_MIN_BLOCKS` selects it. On the 12B, R = 2 unconstrained had drifted to 92 to 95 registers and 2 blocks per SM as the kernel body grew; under a floor of 3 it is 80 registers, 3 blocks and 60.1 tok/s. A floor of 4 caps at 64 and falls back to 54.9, so the budget can be set too tight as well as too loose, and nothing spills at any floor (`LOCAL` is 0 in every instantiation). R = 3 under the same floor of 3 has identical registers and identical residency and moves 17% fewer transactions per output row, and still lands at 58.8. Verified in the shipping build on both checkpoints, generated text byte-identical at `--temp 0` in every arm: | arm | gemma-4-12B-it-4bit | qwen3.8-27B-4bit | |---|---|---| | default (R = 2, floor 3) | 58.3 tok/s | 26.3 tok/s | | R = 3, floor 3 | 57.1 | 26.7 | | floor 1 (no launch bounds) | 54.7 | 26.3 | | R = 1 (stock kernel) | 45.3 | 20.8 | The floor is neutral on the 27B, whose wider projections already fill the machine, and it is worth 6.6% on the 12B. Default is 3 on sm_70 and 1 elsewhere, read from the compute capability the encoder already carries, with `MLXCEL_QMV_MIN_BLOCKS` overriding in both directions. Also tried and removed: batching the scale and bias loads through a warp shuffle. A warp step spans exactly `group_size / elems_per_thread` lanes per group, so four consecutive steps span one group per lane and a single coalesced load plus `__shfl_sync` can serve all four, turning 8 transactions per four steps into 2. It lost in every arm, 60.1 to 51.5 at R = 2 and 58.8 to 46.7 at R = 3, at identical register counts and identical residency, so occupancy explains none of it: the shuffles and the pack and unpack cost more than the transactions they save. It also doubled the instantiation count and pushed the build from 20 minutes to 54, so it is deleted rather than left switched off. Refs #1543
The row-blocked `qmv` issues four loads per step per output row at R = 2, and two of them are the scale and the bias, four bytes of metadata. Removing one by interleaving the pair into a single 32-bit word at load time would cost no runtime arithmetic, unlike the warp shuffle that lost, but it would cost a device-side packing cache keyed by pointer, about 11% more VRAM, and an extension to the `quantized.cpp` overlay. `MLXCEL_QMV_PROBE_NO_SCALE` prices the loads before any of that gets built. The probe reads the scale and the bias from group 0 instead of from the group the step is actually in, so the output is wrong by construction and it is never anything but an A/B instrument. A group-0 read is loop-invariant, so it hoists out of the k-loop and takes both loads with it while leaving every multiply and add where it was. Literal `1` and `0` would have been the obvious spelling and the wrong one: the compiler folds `w_dq * 1 + 0` away and deletes eight HFMA2 per step per row along with the loads, pricing the arithmetic too. Measured on a V100 at 4 bits, four interleaved repeats per arm, decode as the slope between `-n 60` and `-n 200`: | | normal | probe | ceiling | t | |---|---|---|---|---| | gemma-4-12B-it-4bit | 58.4 ± 2.77 | 62.7 ± 3.01 | +7.3% | 2.08 | | qwen3.8-27B-4bit | 26.5 ± 0.79 | 27.7 ± 0.47 | +4.4% | 2.55 | The interleave is not being built. The probe removes both loads for a 7.3% and 4.4% ceiling; interleaving removes one, so the realistic figure is half that, against a VRAM increase and a cache whose lifetime has to be managed. The miss is the useful part. Counting loads says the scale and the bias are half the four a step issues per row, so removing both should have been worth something near the 28% that R = 1 to R = 2 returned for cutting a third of them. It returned 7.3%. The loads are not interchangeable: the scale and the bias are broadcast reads where 32 lanes want 4 values and one transaction serves them from L1, while x and w are 32 distinct addresses. Counting loads is as wrong as counting transactions was, and this is the fourth model in a row to over-predict on this kernel. The probe stays because it is the only instrument this host has: `ncu` returns `ERR_NVGPUCTRPERM` here even under `sudo`, since the container's root holds no `CAP_SYS_ADMIN` and `/proc/driver/nvidia/params` reports `RmProfilingAdminOnly: 1` read-only. Also removes the shared-memory activation staging tried in the previous cycle. Four interleaved repeats put it at +3.9% on the 27B (t = 2.06, and it more than doubled the run-to-run spread, 0.91 against 0.40) and +1.2% on the 12B (t = 0.46), against doubling the instantiation count and taking the build from 23 minutes to 50. That it cut the activation term's global transactions eightfold for almost nothing is what established that this kernel is bound by issued loads rather than by moved bytes. Refs #1543
`rms_norm_small` launches one block per row, so a decode step runs the whole normalization in a single block, and `dispatch_num_chunks` only ever emits 64 or 128 threads for it: `axis_size <= N_READS * 64` takes 64 and every larger case funnels into `dispatch_chunks<128>`. Gemma 4 12B at hidden 3840 gets 128 threads and 4 chunks, Qwen 3.8 27B at 5120 gets 128 and 5. That is four warps on one SM of eighty, each thread walking 32 or 40 elements in a serial loop, for a kernel that runs 336 times per token and measures 5.74 us a launch, 7.7% of the decode window on a V100 for 7.7 KB of work. Upstream's choice is right for prefill, where the grid is already wide and 128-thread blocks pack the machine, which is what the comment above `rms_norm_small` about occupancy is describing. It is wrong when there is one row and nothing else to fill the SM with, so this adds a wide arm that applies only below `kMlxcelWideBlockMaxRows` and leaves the upstream path untouched otherwise. The widths are spaced 128 apart rather than doubling, deliberately. `rms_norm_small`'s fast path needs `axis_size == N_READS * BLOCK * CHUNKS` exactly; miss it and every load takes a bounds check. The 27B hits it today at 128 by 5, and a power-of-two width would have quietly traded that away for the wider block. 640 keeps it, with one chunk instead of five. The 12B has no exact width at or below 1024 other than 480, and 512 leaves it exactly as unaligned as the 128-thread choice it replaces, so nothing is lost there either. `MLXCEL_RMS_WIDE_BLOCK=0` restores upstream's choice for every shape. The wide arm changes how many threads reduce the axis and therefore the last bits of the normalization, so the switch is what an A/B on output quality selects between, and what a rollback uses if a shape regresses. The commit that follows reverts this. It is landed first so that the attempt, its reasoning and its measurements stay in the history rather than disappearing into a discarded working tree: the controlled A/B does not support it, and the reasoning about why the wide arm looked right is worth keeping for whoever reads `dispatch_num_chunks` next and has the same idea. Refs #1543
Reverts 54b4b76. The controlled A/B does not support the change, and the two numbers that first appeared to support it were both artifacts of how they were measured. Interleaved, four repeats per arm, both arms in the same binary through `MLXCEL_RMS_WIDE_BLOCK`: | | wide block | upstream | difference | t | |---|---|---|---|---| | gemma-4-12B-it-4bit | 60.3 ± 2.74 | 59.5 ± 0.87 | +1.3% | 0.52 | | qwen3.8-27B-4bit | 26.6 ± 0.85 | 25.7 ± 0.29 | +3.4% | 1.94 | The first reading was +3.1% and +2.3%, and it compared a run from this build against a baseline measured in the previous one. Nothing guarantees two builds are comparable: this branch has already seen the same unchanged kernel measure 56.0 and 60.1 across builds. The kill switch went in for rollback and caught the error instead, since `MLXCEL_RMS_WIDE_BLOCK=0` then measured 62.8, faster than the arm it was supposed to lose to. The second reading was that the wide arm cut run-to-run standard deviation from 2.77 to 0.25. That was four samples. Measured properly the wide arm sits at 2.74 on the 12B and 0.85 on the 27B, both higher than upstream's 0.87 and 0.29. Reverted rather than left switched off. An 886-line overlay of an upstream file is a resync cost at every MLX pin bump, and +1.3% at t = 0.52 does not earn it. What the negative result establishes is worth more than the change would have been. Going from 128 threads to 640 on the 27B is five times the parallelism with the fast path intact, and it moves throughput 3.4% at best. So `rms_norm_small`'s 5.74 us is not the kernel failing to fill an SM; it is close to what a launch costs on this path at all, and the only thing that moves 336 launches per token is issuing fewer of them. That is fusion in the model graph, not a launch-config choice. Refs #1543
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.
Follows #1678, which has merged, so this is six commits against
maintouchingqmv.cuand the environment variable reference. It replaces #1680, which GitHub closed when #1678's branch was deleted out from under it.What lands
Two changes to
qmv, both on the pre-Ampere decode path, both with an environment override, and both byte-identical to the stock kernel at--temp 0.Row blocking.
qmvgives each warp one output row and has that warp read the whole activation vector, so a launch movesn * k * 2bytes of activations againstn * k / 2bytes of 4-bit weights: four bytes of activation for every byte of weight. Giving each warpRconsecutive output rows and loading the activation tile once for all of them divides that term byRand leaves weight traffic alone. The per-row arithmetic is untouched, which is what keeps the result bit-identical.A pinned resident-block floor. The R sweep turned out to be a curve in occupancy rather than in loads, and it tracks the register column exactly: R = 1 at 40-48 registers and 4-5 blocks per SM gives 44.7 tok/s, R = 2 at 64-72 and 3-4 gives 56.0, R = 3 at 80-110 and 2-3 gives 53.4, R = 4 at 96-142 and 1-2 gives 45.2, R = 6 at 128 and 2 gives 42.7. Transactions per output row fall at every step and past R = 2 every step also costs a resident block, and the occupancy wins each time. So
__launch_bounds__(256, min_blocks)hands ptxas the register budget directly instead of trying to reach it through the row count.Measured on a V100 at 4 bits, decode as the slope between
-n 60and-n 200, all defaults in place, generated text byte-identical in every arm:Defaults are R = 2 with a floor of 3, on sm_70 only, read from the compute capability the encoder already carries. Turing stays at R = 1 with no floor: it has a different register file per SM, no sm_75 device was available to measure, and the R = 4 reversal is a direct demonstration that this tradeoff does not extrapolate.
MLXCEL_QMV_ROWS_PER_WARP=1is the rollback.What is here as an instrument
MLXCEL_QMV_PROBE_NO_SCALEreads the scale and the bias from group 0 rather than the group the step is in, so the output is wrong by construction. It exists because two of the four loads a step issues per output row are the scale and the bias, four bytes of metadata, and interleaving them into a single 32-bit word at load time would have cost a device-side packing cache, about 11% more VRAM, and an extension to thequantized.cppoverlay. The probe prices the loads before any of that: removing both is worth 7.3% on the 12B and 4.4% on the 27B, so the interleave that removes one is not being built.The probe stays because
ncucannot run here. The box is a Docker container whose root holds noCAP_SYS_ADMIN,/proc/driver/nvidia/paramsis read only, and it reportsRmProfilingAdminOnly: 1, so Nsight Compute returnsERR_NVGPUCTRPERMeven undersudo. Nsight Systems works, but only with--trace=cuda: the default trace set breakscuModuleLoadDataExon MLX's JIT PTX with an empty JIT log, which looks like a corrupt cache and is not one.What was tried and rejected
Each of these is worth a line because the next person to read
qmv.cuwill have the same ideas.__restrict__on the row-blocked kernel. No change to register counts, no change to throughput.__shfl_synccan serve all four and turn 8 transactions per four steps into 2. It lost everywhere: 60.1 to 51.5 at R = 2 and 58.8 to 46.7 at R = 3, at identical register counts and identical residency. The shuffles and the pack and unpack cost more than the transactions they save.rms_normblock when the grid is one row. Landed and reverted in this branch on purpose, so the attempt and its numbers stay in the history: see the revert commit.rms_norm_smallis 7.7% of the decode window at 336 launches per token, and upstream only ever gives it 64 or 128 threads, which is four warps on one SM of eighty. Going to 640 threads on the 27B with the aligned fast path intact moves throughput 3.4% at t = 1.94, and the 12B 1.3% at t = 0.52.The staging result is the useful one: cutting the activation term's global transactions eightfold for almost nothing is what established that this kernel is bound by the load instructions its warps issue rather than by the bytes they move.
Not covered
sm_75 is untested, by choice, and defaults to the stock path.
Decode targets are not met. The 12B is at 60 tok/s against a 110 target and the 27B at 27 against 50.
qmvis 68.6% of the decode window, so halving it yields 34% of the 45% the 12B target needs; the targets requireqmvroughly three times faster. One structural idea remains unexplored: rewriting the dot product assum_g [ scale * sum(x * w_int) + bias * sum(x) ]accumulates the raw integer product per group and applies the scale and bias once per group, which turns the per-row accumulator from sixteen elements into one scalar and reopens the register budget that pins R at 2. It changes the summation order, so the byte-identity check that validated everything in this PR stops applying and a perplexity gate has to replace it.Refs #1543