Summary
mx.fast.metal_kernel custom kernels are silently dispatched over a truncated threadgrid whenever a grid dimension is smaller than the matching requested threadgroup dimension — which is the common case (e.g. one threadgroup per query row with a 256-wide threadgroup).
Root cause
In CustomKernel::eval_gpu (mlx/backend/metal/custom_kernel.cpp):
MTL::Size group_dims =
MTL::Size(std::min(tx, gx), std::min(ty, gy), std::min(tz, gz));
MTL::Size grid_dims = MTL::Size(gx, gy, gz);
compute_encoder.dispatch_threads(grid_dims, group_dims);
dispatch_threads interprets grid_dims as thread counts, so with group_dims clamped to min(tx, gx) etc. the effective threadgroup count per dimension is ceil(min(tx,gx)/tx_per_group) — a small prefix of the requested grid. For a per-row kernel with requested grid (24, Q, 1) and threadgroup (256, 1, 1), the effective dispatch covers 24 threads of x per row (1 threadgroup of 24 threads) instead of 256 — i.e. ~1/256 of the requested work executes.
Behavior
- The truncation is stable across repeated same-geometry dispatches (verified: 20+ dispatches all truncated at a fresh geometry; re-dispatch does not clear it).
- It is invisible to every check except full coverage of the requested grid (e.g. a probe kernel writing each thread's unique index and comparing the output to
{1..N}).
- Benchmark impact: a custom kernel's timing measured through this path can appear orders of magnitude faster than the true full-work cost. In our case a flash-SDPA kernel measured 3.5–4.8× faster than the incumbent with truncated dispatch, and 57.7× slower at true full work. Any custom-kernel performance claim made before this fix is invalid without per-dispatch fullness evidence.
Repro (minimal, no MLXLM, no attention code)
import MLX
// probe kernel: out[idx] = idx for each executing thread
let src = """
kernel void full_probe(device const int2& in [[buffer(0)]],
device int& out [[buffer(1)]]) {
long idx = (long)(in.y * in.x + threadgroup_position_in_grid.x) * 256L
+ position_in_threadgroup.x;
out[idx] = (int)idx;
}
"""
let fn = MLXFast.metalKernel(
name: "full_probe",
inputNames: ["in"],
outputNames: ["out"],
source: src)
let inArr = MLXFast.metalKernel(/* ... */ ) // unused values, shape {gy*gx, 2} int32
let out = MLX.zeros(shape: [24 * 64 * 256], dtype: MLX.DType.int32)!
// grid (24, 64, 1), threadgroup (256, 1, 1)
let r = try fn(inputs: [inArr], outputShapes: [out.shape], outputDTypes: [.int32],
threadsPerThreadgroup: (256, 1, 1), threadsPerGrid: (24, 64, 1))
MLX.eval(r.0)
// Full coverage: out == arange(1, N+1). Before the fix: only a small
// non-zero prefix (truncated dispatch); stable across repeated calls.
C++ equivalent: fast::metal_kernel(...) in tests/gpu_tests.cpp (added in the linked PR).
Evidence from our investigation (Qwen3.8 MTP server, Qwen3.8-27B, M5 Pro)
- Trivial probe (no attention code): dispatch alternates between RECOMBINED (24 groups × 64 threads) and FULL (1536 groups × 256 threads) on every other dispatch; pure same-geometry sequences stay truncated for all dispatches (20/20).
- Q sweep: deterministic across 3 runs; truncation fires for all Q (not just small Q), first dispatch of a fresh geometry included.
- Production geometry: Q=2048, prefix=8192 (our prefill chunk): truncated dispatch writes 49,152 elements (= 24 × 2048) instead of 12,582,912 (= 24 × 2048 × 256).
- Pins at time of report: mlx-swift
2bebe4e9ad127758ebcd76c6ad45a1740d0d2852, mlx C++ 1f8e74e3f12f31365464a6867c6579f0e9b29d85 (2026-08-25).
- Full-work re-bench after local fix: flash kernel is bit-exact deterministic and correct (max|diff| 6.1e-5 vs bf16 reference) at true full work, confirming the truncation was the sole cause of the earlier measurements.
Fix
Dispatch the requested threadgroup dimensions unclamped through dispatch_threadgroups:
MTL::Size group_dims = MTL::Size(tx, ty, tz);
MTL::Size grid_dims = MTL::Size(gx, gy, gz);
compute_encoder.dispatch_threadgroups(grid_dims, group_dims);
(See linked PR for the patch + a full-coverage regression test in tests/gpu_tests.cpp.)
Versions
- MLX:
1f8e74e3f12f31365464a6867c6579f0e9b29d85 (2026-08-25)
- mlx-swift:
2bebe4e9ad127758ebcd76c6ad45a1740d0d2852
- macOS 26.5.2 (25F84), Apple M5 Pro, 48 GB
Summary
mx.fast.metal_kernelcustom kernels are silently dispatched over a truncated threadgrid whenever a grid dimension is smaller than the matching requested threadgroup dimension — which is the common case (e.g. one threadgroup per query row with a 256-wide threadgroup).Root cause
In
CustomKernel::eval_gpu(mlx/backend/metal/custom_kernel.cpp):dispatch_threadsinterpretsgrid_dimsas thread counts, so withgroup_dimsclamped tomin(tx, gx)etc. the effective threadgroup count per dimension isceil(min(tx,gx)/tx_per_group)— a small prefix of the requested grid. For a per-row kernel with requested grid(24, Q, 1)and threadgroup(256, 1, 1), the effective dispatch covers24threads of x per row (1 threadgroup of 24 threads) instead of 256 — i.e. ~1/256 of the requested work executes.Behavior
{1..N}).Repro (minimal, no MLXLM, no attention code)
C++ equivalent:
fast::metal_kernel(...)intests/gpu_tests.cpp(added in the linked PR).Evidence from our investigation (Qwen3.8 MTP server, Qwen3.8-27B, M5 Pro)
2bebe4e9ad127758ebcd76c6ad45a1740d0d2852, mlx C++1f8e74e3f12f31365464a6867c6579f0e9b29d85(2026-08-25).Fix
Dispatch the requested threadgroup dimensions unclamped through
dispatch_threadgroups:(See linked PR for the patch + a full-coverage regression test in
tests/gpu_tests.cpp.)Versions
1f8e74e3f12f31365464a6867c6579f0e9b29d85(2026-08-25)2bebe4e9ad127758ebcd76c6ad45a1740d0d2852