From 346eff7503a22cfe2e020cf595260ce347961af4 Mon Sep 17 00:00:00 2001 From: Chris Wong Date: Sat, 19 Sep 2026 17:55:46 +0100 Subject: [PATCH] Fix CustomKernel Metal dispatch silently truncating the threadgrid CustomKernel::eval_gpu computed group_dims as min(threadgroup, grid) per dimension and dispatched through dispatch_threads. dispatch_threads interprets its grid argument as thread counts, so the effective threadgroup count became ceil(min(tx,gx)/max_tx_per_threadgroup), a small prefix of the requested grid whenever a grid dimension is smaller than the matching threadgroup dimension (the common case). The truncation is stable across repeated same-geometry dispatches, so re-dispatching cannot clear it; only a full-coverage check of the requested grid detects it. Use dispatch_threadgroups with the requested threadgroup dimensions unclamped, and add a regression test that dispatches a probe kernel over several (grid, threadgroup) geometries - including the production per-row geometry with grid.x < threadgroup width - and verifies exact coverage of every requested thread on the first dispatch and on a repeated same-geometry dispatch. --- mlx/backend/metal/custom_kernel.cpp | 11 ++-- tests/gpu_tests.cpp | 83 +++++++++++++++++++++++++++++ 2 files changed, 91 insertions(+), 3 deletions(-) diff --git a/mlx/backend/metal/custom_kernel.cpp b/mlx/backend/metal/custom_kernel.cpp index b73dd72def..a965791d2c 100644 --- a/mlx/backend/metal/custom_kernel.cpp +++ b/mlx/backend/metal/custom_kernel.cpp @@ -93,10 +93,15 @@ void CustomKernel::eval_gpu( } const auto [gx, gy, gz] = grid_; - MTL::Size group_dims = - MTL::Size(std::min(tx, gx), std::min(ty, gy), std::min(tz, gz)); + // dispatch_threadgroups takes the threadgroup dimensions as-is; using + // dispatch_threads here would treat the requested *threadgroup* sizes as + // thread counts and silently truncate the dispatch to + // ceil(thread_count / max_threads_per_threadgroup) threadgroups, which is + // fewer than requested whenever a dimension of grid_ is smaller than the + // matching threadgroup dimension (the common case). + MTL::Size group_dims = MTL::Size(tx, ty, tz); MTL::Size grid_dims = MTL::Size(gx, gy, gz); - compute_encoder.dispatch_threads(grid_dims, group_dims); + compute_encoder.dispatch_threadgroups(grid_dims, group_dims); compute_encoder.add_temporaries(std::move(copies)); } diff --git a/tests/gpu_tests.cpp b/tests/gpu_tests.cpp index 8bef07a616..34bb39f096 100644 --- a/tests/gpu_tests.cpp +++ b/tests/gpu_tests.cpp @@ -712,3 +712,86 @@ TEST_CASE("test layer norm vjp bias grad race") { } CHECK(worst <= 1e-5); } + +TEST_CASE("fast metal kernel dispatches the full requested grid") { + if (default_device().type != Device::gpu) { + return; + } + + // Regression test for a CustomKernel::eval_gpu bug that clamped the + // requested threadgroup dimensions against the grid dimensions + // (group_dims = (min(tx,gx), min(ty,gy), min(tz,gz))) and dispatched the + // result through dispatch_threads. The combined effect was a silently + // truncated dispatch: whenever a grid dimension was smaller than the + // matching threadgroup dimension (the common case, e.g. a per-row grid + // with a 256-wide threadgroup), only a small prefix of the requested + // threadgroups ever executed. The truncation was stable across repeated + // same-geometry dispatches, so re-dispatching could not clear it; the only + // reliable detector is full thread coverage of the requested grid. + // It also silently truncated work on every custom-kernel dispatch in the + // affected window, making custom-kernel timing measurements appear + // dramatically faster than the true full-work cost: a timing claim for a + // custom kernel is not valid without per-dispatch fullness evidence. + // + // The probe kernel writes each executing thread's unique index + // ((grid.y * grid.x + group.x) * tx + thread.x, with grid.z = tz = 1) + // into its own output slot; full coverage of the requested gx * ty * tx + // threads is equivalent to the output being exactly {1, 2, ..., N}. + + const std::string src = R"( +kernel void full_probe(device const int2& pos [[buffer(0)]], + device int& out [[buffer(1)]]) { + long idx = (long)(pos.y * pos.x + threadgroup_position_in_grid.x) * 256L + + position_in_threadgroup.x; + out[idx] = (int)idx; +} +)"; + + auto make_kernel = [&]() { + return fast::metal_kernel( + "full_probe", + {"in"}, + {"out"}, + src, + "", + true, + false); + }; + + // (grid, threadgroup) pairs. The first pair is the production geometry + // that motivated the test: one threadgroup per query row, a 256-wide + // threadgroup, grid.x (24) far below the threadgroup width. + const std::vector> grids = { {24, 64, 1}, + {64, 24, 1}, + {24, 2048, 1} }; + const std::tuple tg = {256, 1, 1}; + + for (auto [gx, gy, gz] : grids) { + auto in = array(ones({gy * gx, 2}, int32, Device::gpu)); + const long N = (long)gx * gy * std::get<0>(tg); + Shape out_shape; + out_shape.push_back((int32_t)N); + auto ref = copy(arange(1, N + 1, int32, Device::cpu), Device::gpu); + + auto fn = make_kernel(); + + // First dispatch at a fresh geometry, then a repeated same-geometry + // dispatch: before the fix the truncation was stable across the whole + // same-geometry sequence, so both must show full coverage. + for (int iter = 0; iter < 2; ++iter) { + auto result = fn( + {in}, + {out_shape}, + {int32}, + {gx, gy, gz}, + tg, + {}, + {}, + false, + {}); + auto res = result[0]; + eval(res); + CHECK(array_equal(res, ref, Device::gpu).item()); + } + } +}