<!-- llms-explorer concept facts · https://llms-explorer.com/tree/mlx-gemv-wide-bf16-fp16-few-row-matmul-kernel-pr/ · pack 2026-10-05 · ~1935 tokens -->

# MLX gemv_wide bf16 fp16 few-row matmul kernel PR 3888

> 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...

Parent: [Mac local LLMs: MLX kernels, numerics and internals](https://llms-explorer.com/tree/mac-local-llms-mlx-kernels-numerics-and-internals/) · 1 facets · 32 facts · page: https://llms-explorer.com/tree/mlx-gemv-wide-bf16-fp16-few-row-matmul-kernel-pr/

## 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: `asserted`
- 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: `asserted`
- A register tile holds at most five vectors because wider tiles collapse occupancy (source comment). Each extra pass re-streams the matrix. — source: `asserted`
- 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: `asserted`
- 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: `asserted`
- 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: `asserted`
- 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: `asserted`
- 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: `asserted`
- 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: `asserted`
- 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: `asserted`
- 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: `asserted`
- 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: `asserted`
- 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: `asserted`
- 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: `asserted`
- 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: `asserted`
- Whether fp32 or fp8 inputs will get a wide route. — source: `asserted`
- 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: `asserted`
- `gemv_wide_config` returns no plan when architecture generation is below 15. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/matmul.cpp)
- `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](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/matmul.cpp)
- The wide route computes `passes = (M + 4) / 5` and declines when passes exceeds 3, so M is capped at 15. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/matmul.cpp)
- `vecs_per_tg = (M + passes - 1) / passes`, so rows are balanced across passes. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/matmul.cpp)
- `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](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/matmul.cpp)
- The kernel launches with threadgroup (32, k_lanes/8, 1) and grid (grid_x, ceil(N/4), batch). — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/matmul.cpp)
- The kernel name is `gemv_wide_<dtype>_nv<vecs_per_tg>_kl<k_lanes>` with `has_batch` and `do_axpby` function constants. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/matmul.cpp)
- 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](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/matmul.cpp)
- `gather_mm_wide` serves GatherMM through `gemv_wide_gather` only when b is a transposed view, and declines a transposed a. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/matmul.cpp)
- A source comment states pre-M3 generations are limited by load issue rate rather than bandwidth and keep the existing kernels. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/matmul.cpp)
- 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](https://github.com/ml-explore/mlx/pull/3888)
- PR 3888 benchmark shapes come from a Qwen3.6 deployment: the layers a quantized model keeps in bf16. — [source](https://github.com/ml-explore/mlx/pull/3888)
- `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](https://github.com/ml-explore/mlx/pull/3888)
- `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](https://github.com/ml-explore/mlx/pull/3888)
- In a maintainer's chained benchmark (D=4096) fp16 x@W.T cost 0.090-0.097 ms for N=8-14 and 0.339 ms at N=16. — [source](https://github.com/ml-explore/mlx/issues/3086)
