MLX gemv_wide bf16 fp16 few-row matmul kernel PR 3888
Parent: Mac local LLMs: MLX kernels, numerics and internals · Published reference · snapshot 2026-10-05
↓ Facts as markdownall context files
Eligibility (`gemv_wide_config`, main): architecture generation >= 15; M >= 2 and N >= 2; K a multiple of 4; output dtype float16 or bfloat16; vec4-aligned operands (matrix and vector offsets multiples of 8, leading dimensions multiples of 4, all batch strides multiples of 4). Otherwise the exist...
These notes link each claim to its source. A source may be a research report hosted on this site rather than the primary document. A published reference means the content is available; it does not certify independent review or accuracy.Read the editorial policy and follow the sources before relying on a claim.
Facts
- Eligibility (`gemv_wide_config`, main): architecture generation >= 15; M >= 2 and N >= 2; K a multiple of 4; output dtype float16 or bfloat16; vec4-aligned operands (matrix and vector offsets multiples of 8, leading dimensions multiples of 4, all batch strides multiples of 4). Otherwise the existing gemv or GEMM routes run. [source]
- Pass count: `passes = (M + 4) / 5`; if passes > 3 the config declines. This is why the ceiling is M = 15, not a tuned crossover: the source comment says a fourth pass no longer pays. Rows are balanced: `vecs_per_tg = ceil(M / passes)`, so M=6 runs two passes of 3 vectors, M=11 three passes of 4. [source]
- A register tile holds at most five vectors because wider tiles collapse occupancy (source comment). Each extra pass re-streams the matrix. [source]
- Launch shape: threadgroup (32, k_lanes/8, 1); grid (grid_x, ceil(N/4), batch). `k_lanes` is 32 when passes == 1 or N <= 64, else 16 (full simdgroup rows double threads per group and halve each lane's K share where the grid is thinnest). `grid_x = 1` for N >= 65536 (vocab-wide outputs; chunks beyond the grid loop inside one threadgroup so the matrix streams once), else `grid_x = passes`. [source]
- Kernel name is `gemv_wide_<dtype>_nv<vecs_per_tg>_kl<k_lanes>` with function constants `has_batch` and `do_axpby`; AddMM reaches it with alpha/beta bias handling. A gather variant `gemv_wide_gather_...` serves GatherMM when b is a transposed view; it declines a transposed `a` and copies `a` to contiguous when neither layout matches. [source]
- Routing precedence: in both `matmul` and `addmm` the wide route is tried in the gemv specialization section, ahead of the split-K and NAX GEMM selection, and only when `a` is not transposed and `b` is transposed. M == 1 stays on plain gemv. [source]
- Stated rationale for the gen-15 gate: pre-M3 generations are limited by load issue rate rather than bandwidth and do not profit from the amortized weight stream. [source]
- PR 3888 opened 2026-07-21 by jessegross with shapes from a Qwen3.6 deployment (the layers a quantized model keeps in bf16); angeloskath approved 2026-07-22 after adding a GPU-requirement check for the wide matmul tests and a lint fix. [source]
- A visible cliff the kernel sits under was in issue 3086 (2026-01-31, M2 Max, mlx 0.30.3): in a maintainer's chained benchmark fp16 matmul jumps from 0.097 ms at N=14 to 0.339 ms at N=16 (D=4096) because the GEMM tile path starts there. Hardware for the maintainer's chained numbers was not stated. [source]
- Eligibility is dtype-strict: an fp32 activation or output never takes gemv_wide, so fp32 few-row matmuls keep the old routes (and, on gen-17, the TF32 NAX GEMM for GEMM shapes). [source]
- Misaligned slices decline silently: a K not divisible by 4, an odd leading dimension, or an unaligned view falls back, so shapes that should win may not. [source]
- For N <= 64 and M <= 15 (tiny projections such as a 32-wide `in_proj_a/b`) the kernel gets the full-simdgroup config; the 6.0x M5 Max figure in the PR is for exactly that case and should not be generalized to wide matrices (lm_head gains are 1.3-2.3x). [source]
- M5 gain pattern differs from M3/M4: tiny-N shapes gain most on M5 Max (4.7-6.0x vs 1.7-3.0x) while the vocab-wide lm_head gains least on M5 Max (1.3-1.4x vs 1.7-2.3x), consistent with the M5 GPU already being closer to bandwidth-bound on the wide case. This is an inference from the PR table. [source]
- None found; the PR has no recorded objection. Open numeric gap: the PR reports kernel time with uncached weights, not end-to-end tokens/s. [source]
- Whether the cap of five vectors and three passes is re-tuned for M5-class chips; the source gates only on generation >= 15, not on M5. [source]
- Whether fp32 or fp8 inputs will get a wide route. [source]
- End-to-end effect on a hybrid model's decode and MTP verify (the in_proj and lm_head layers); not reported in the PR. [source]
- `gemv_wide_config` returns no plan when architecture generation is below 15. [source]
- `gemv_wide_config` declines when M <= 1, N <= 1, K is not a multiple of 4, the output dtype is not float16 or bfloat16, or operands are not vec4-aligned. [source]
- The wide route computes `passes = (M + 4) / 5` and declines when passes exceeds 3, so M is capped at 15. [source]
- `vecs_per_tg = (M + passes - 1) / passes`, so rows are balanced across passes. [source]
- `k_lanes` is 32 when passes is 1 or N <= 64, else 16; `grid_x` is 1 when N >= 65536, else the pass count. [source]
- The kernel launches with threadgroup (32, k_lanes/8, 1) and grid (grid_x, ceil(N/4), batch). [source]
- The kernel name is `gemv_wide_<dtype>_nv<vecs_per_tg>_kl<k_lanes>` with `has_batch` and `do_axpby` function constants. [source]
- The matmul and addmm paths call `gemv_wide` when `a` is not transposed and `b` is transposed, before split-K and NAX GEMM dispatch. [source]
- `gather_mm_wide` serves GatherMM through `gemv_wide_gather` only when b is a transposed view, and declines a transposed a. [source]
- A source comment states pre-M3 generations are limited by load issue rate rather than bandwidth and keep the existing kernels. [source]
- PR 3888 was opened 2026-07-21 and approved by angeloskath on 2026-07-22 after he added a GPU requirement check for the wide matmul tests. [source]
- PR 3888 benchmark shapes come from a Qwen3.6 deployment: the layers a quantized model keeps in bf16. [source]
- `lm_head` [M,2048]x[2048,248320] speedups at M=2/4/8 are 1.7/1.7/1.8x (M3 Ultra), 2.3/2.2/2.1x (M4 Pro), 1.4/1.4/1.3x (M5 Max). [source]
- `in_proj_a/b` [M,2048]x[2048,32] speedups at M=2/4/8 are 3.0/2.3/2.3x (M3 Ultra), 2.3/1.8/1.7x (M4 Pro), 6.0/4.7/4.8x (M5 Max). [source]
- In a maintainer's chained benchmark (D=4096) fp16 [email protected] cost 0.090-0.097 ms for N=8-14 and 0.339 ms at N=16. [source]
Children
- No children recorded.