Repository navigation
Bound winograd conv2d working set by tiling the batch - #4102
Conversation
52fb201 to
413078a
Compare
|
Verified on an M4 Pro Mac mini with 48 GB, macOS 26.5 — the reporter's Control is Failure starts at batch 40, the number @kayadibi1 reported on 48 GB. Peak Also on an M5 Max via the injected budget: at 512 MB peak drops 1.10 -> 0.36 GB No opinion on the tile-versus-fallback heuristic, which I have not stressed |
|
|
||
| size_t itemsize = in.itemsize(); | ||
| int padded_h = | ||
| 6 * ((conv_params.iS[0] + 2 * conv_params.pad[0] - 2 + 5) / 6) + 2; |
There was a problem hiding this comment.
Can you rebase and fix the conflict with https://github.com/ml-explore/mlx/pull/4258/changes#diff-d89046a37f9b2d939417b57314e679e1d9aa2f64a00b0a45b7e4d9656a75056dR903?
| size_t working_set = d.mtl_device()->recommendedMaxWorkingSetSize(); | ||
| // Test hook: pretend the GPU can only keep this many bytes resident, so the | ||
| // budget arithmetic below can be exercised at sizes that fit in CI. | ||
| if (auto env_ws = env::get_var("MLX_CONV_WINOGRAD_WORKING_SET", ""); |
There was a problem hiding this comment.
I think you can just do
| if (auto env_ws = env::get_var("MLX_CONV_WINOGRAD_WORKING_SET", ""); | |
| if (int env_ws = env::get_var("MLX_CONV_WINOGRAD_WORKING_SET", 0); |
| // but charging all live allocations only errs toward a smaller tile. The | ||
| // input, weights and output are already allocated and inside this figure, | ||
| // so only winograd's scratch is left to make room for. | ||
| size_t used = get_active_memory() + filt_bytes; |
There was a problem hiding this comment.
We don't dispatch kernels depending on current memory, it is not reliable and make behavior unpredictable.
| # Get the previous run's temporaries off the books first. | ||
| mx.synchronize() | ||
| mx.reset_peak_memory() | ||
| base = mx.get_active_memory() |
There was a problem hiding this comment.
I would prefer dropping the memory measurements in tests, empirically they tend to be very flaky.
413078a to
3772dec
Compare
Winograd's scratch is a multiple of the input, all referenced by one command buffer; once that outgrows the GPU's recommended working set the kernels silently produce an all-zero output (issue ml-explore#3979). Size a batch tile from the available working set, reuse one set of scratch buffers across tiles, and fall back to the scratch-free implicit gemm when not even one element fits or a tile would carry too little gemm work to amortize its fixed cost. The real threshold needs tens of GB, so tests shrink the budget instead: MLX_CONV_WINOGRAD_WORKING_SET injects the byte budget and MLX_CONV_WINOGRAD_TILE_BATCH forces a tile size, exercising uniform and short final tiles, the automatic selector, and both fallbacks at CI sizes.
Size the budget from the conv's own tensors instead of active memory, read the working-set override with env::get_var like other env vars, and pin each test run to its path through bit-equality with the untiled result instead of peak-memory measurements.
3772dec to
880e6bd
Compare
* Return tuple in meshgrid (ml-explore#4229) * Add endpoint parameter to linspace (ml-explore#4184) Co-authored-by: Cheng <git@zcbenz.com> * Fix vmap of partition/argpartition dropping the kth argument (ml-explore#4116) * Fix nan_to_num replacing inf with 0 for float16 and bfloat16 (ml-explore#4222) Co-authored-by: codeAnqiang-ma <273298913+codeAnqiang-ma@users.noreply.github.com> Co-authored-by: Cheng <git@zcbenz.com> * Fix einsum not broadcasting batch dimensions in batched tensordot (ml-explore#4125) Co-authored-by: Cheng <git@zcbenz.com> * Dequantize in float32 (ml-explore#4241) * chore: Reject complex in erf and erfinv (ml-explore#4243) * Fix cpu compilation failure of abs with uint (ml-explore#4240) Co-authored-by: Cheng <git@zcbenz.com> * Fix quantize matrix multiplication floor issue (ml-explore#4251) * Only use MPI backend for world size > 1 (ml-explore#4210) * chore: Reject complex in expm1, sigmoid and arctan2 (ml-explore#4257) * Decompose small kernel-depth 3D convs into 2D convs (ml-explore#3785) Co-authored-by: katlun-lgtm <264247399+katlun-lgtm@users.noreply.github.com> Co-authored-by: Cheng <git@zcbenz.com> * Fix Metal sort of a view with a negative stride (ml-explore#4252) * Mirror the depth axis in the decomposed 3D conv when flipped (ml-explore#4277) * Fix Metal row reductions on negative-stride views (ml-explore#4267) Co-authored-by: Fu Xiaonan <214359569+FU-max-boop@users.noreply.github.com> * [CUDA] Fix custom kernel cache collision for same name, different source (ml-explore#4273) Co-authored-by: Cheng <git@zcbenz.com> * Fix ops rejecting integers larger than INT32_MAX (ml-explore#4255) Co-authored-by: Feli <feli@hnu.edu.cn> Co-authored-by: Cheng <git@zcbenz.com> * Fix var/std for complex numbers (ml-explore#4260) * Fix int32 overflow in conv padded input and pad shapes (ml-explore#4258) Co-authored-by: Cheng <git@zcbenz.com> * chore: Reject complex in remainder (ml-explore#4270) * chore: Compare the macOS SDK version as a version when gating JACCL (ml-explore#4286) * Clamp ring socket transfers so a payload of 2 GiB or more can be sent (ml-explore#4281) Co-authored-by: Cheng <git@zcbenz.com> * chore: Use normalize_axis_index in split/unstack/partition/topk (ml-explore#4288) * Remove grouped output in CI (ml-explore#4195) * [CUDA] Fix finding cuda 13 headers in JIT compilation (ml-explore#3995) * Refactor wheel building script (ml-explore#3818) * Make mx.compile cache erasing thread safe (ml-explore#4248) Co-authored-by: yentur <mr.yentur@gmail.com> * Add builds for free-threaded python (ml-explore#3812) * Fix int32 overflow in concatenate/repeat/kron (ml-explore#4303) * python: Widen list elements that do not fit in int32 to int64 (ml-explore#4305) * Propagate CPU errors to events (ml-explore#3742) Co-authored-by: Alessio Pollero <alessio.pollero@gmail.com> * Fix mx.arange dtype inference overflow regression (ml-explore#4324) * Add workflow to update pull request limit bypass list (ml-explore#4320) * Support head dimension 72 in Metal full attention (ml-explore#4330) * Patch bump to 0.32.2 (ml-explore#4333) * Preserve subnormal float values when casting to bool (ml-explore#4224) * python: Support assigning through a bare Ellipsis index (ml-explore#4314) * Fix divmod truncating the quotient for floats (ml-explore#4108) Co-authored-by: Cheng <git@zcbenz.com> * Add force_fused option to scaled_dot_product_attention (ml-explore#4185) * chore: Reject negative eps in the normalization layers (ml-explore#4312) * Bound GGUF metadata string/array values against the file mapping (ml-explore#4212) Co-authored-by: x14ngch3n <x14ngch3n@users.noreply.github.com> Co-authored-by: Cheng <git@zcbenz.com> * Read each K/V byte once in gqa-8 decode attention (ml-explore#4077) * Fix fft vmap and jvp for transforms over a subset of axes (ml-explore#4138) * Fix median dropping NaN (ml-explore#4146) * Fix the CPU scan over a size one axis with a padded stride (ml-explore#4139) Co-authored-by: Cheng <git@zcbenz.com> * chore: Validate the optimizer betas at construction (ml-explore#4310) Co-authored-by: Cheng <git@zcbenz.com> * `RMSNormVJP` backward writes a full `{n_rows, D}` `gw_temp` intermediate (ml-explore#4293) * [Bug]: add default none value to axis parameter of the take_along_axis (ml-explore#4357) Co-authored-by: Anastasiia Filippova <a_filippova@apple.com> * Add a fused full-attention path for head_dim 256 on NAX devices (ml-explore#3842) Co-authored-by: Cheng <git@zcbenz.com> * Update nanobind to 2.15.0 (ml-explore#4337) * Skip unnecessary simdgroup computations for quantised MOE matmuls on NAX (ml-explore#4352) * Add AI usage policy (ml-explore#4331) Co-authored-by: Jake Bowhay <60778417+j-bowhay@users.noreply.github.com> * Raise cpu stream errors from synchronize (ml-explore#4338) Co-authored-by: Cheng <git@zcbenz.com> * chore: Validate eps in Adam at construction (ml-explore#4361) Co-authored-by: Anastasiia Filippova <a_filippova@apple.com> * Bound winograd conv2d working set by tiling the batch (ml-explore#4102) Co-authored-by: Cheng <git@zcbenz.com> * Use a 32-row block in qmm_t_nax when one block covers all of M (ml-explore#4171) * chore: Deduplicate fftshift and ifftshift (ml-explore#4318) * Fix Log and Equal is_equivalent ignoring primitive state (ml-explore#4266) Co-authored-by: Cheng <git@zcbenz.com> * Stabilize reduced-precision InstanceNorm (ml-explore#4230) * chore: Normalize negative axes in sort and argsort (ml-explore#4332) * Clean up main thread compile cache before python interpreter shuts down (ml-explore#4373) * chore: Check malformed jaccl hostfile that miss rdma in pairs (ml-explore#4284) Co-authored-by: Cheng <git@zcbenz.com> * Round mxfp8 block scales up to avoid saturation (ml-explore#4353) Co-authored-by: Daniel Hiltgen <daniel.hiltgen@ollama.com> Co-authored-by: Cheng <git@zcbenz.com> * Add support for the __array_namespace_info__ (ml-explore#4334) * Stop a failed CUDA graph commit from poisoning the encoder (ml-explore#4356) Co-authored-by: Cheng <git@zcbenz.com> * Fix quantized kernels in JIT build (ml-explore#4372) Co-authored-by: Cheng <git@zcbenz.com> * Avoid zero work in stride-2 ConvTranspose3d (ml-explore#4343) * [CUDA] Ce fused kernel (ml-explore#3947) * Fix cpu exclusive scan for complex numbers (ml-explore#4272) Co-authored-by: Cheng <git@zcbenz.com> * Support Relocatable CUDA DLLs on Windows (ml-explore#4382) * Use cast_to for fused AsType in compiled Metal kernels (ml-explore#4351) Co-authored-by: katlun-lgtm <katlun@windyviews.com> Co-authored-by: Cheng <zcbenz@gmail.com> * python: Declare DLPackCompatible protocol members as methods (ml-explore#4384) * Fix quantizing sliced arrays (ml-explore#4381) * Fix einsum dropping a trailing empty subscript (ml-explore#4299) Co-authored-by: Cheng <git@zcbenz.com> * Add script to run python tests (ml-explore#4393) * Hold GIL in AttachedData destructor (ml-explore#4391) * Bound Metal buffer COUNT, not just bytes, in MetalAllocator The Metal allocator throws `[metal::malloc] Resource limit (N) exceeded` when num_resources_ (the live+cached Metal buffer COUNT) reaches resource_limit_ (the iogpu.rsrc_limit sysctl, default ~499000). Freed buffers are recycled into a size-keyed cache whose only trim is by BYTES (release_cached_buffers takes a bytes-to-free target, max_pool_size_ ~= physical RAM). Under churn with many distinct buffer shapes (varied prompt lengths, growing KV caches, multiple co-resident models) the cache fills with entries never reused at that exact size, so the COUNT climbs to the limit while byte usage stays modest and the byte trim never fires — the process crashes mid-inference on a machine with most of its RAM free. malloc() now also reclaims by count: when num_resources_ crosses a 90% high-water mark of resource_limit_, it clears the (pure-reuse) buffer cache so the count drops back to the live working set. Clearing the cache only costs re-allocation, never correctness, so the count limit becomes unreachable by any request mix or batching method while the existing byte limits keep total memory bounded. Adds get_num_resources()/get_resource_limit() to the public memory API (metal + no_gpu + cuda backends) so the count and its ceiling are observable from callers. Adds an MLX_RESOURCE_LIMIT env override that can only LOWER the ceiling (clamped to the OS limit, strictly validated) to exercise the trim deterministically and as an operator safety valve. * perf(mlx): opt-in Gemma 4 expert-QMM tile kernel with parallel descriptor builder (#4) * perf(mlx): add opt-in Gemma 4 expert-QMM tile kernel with parallel descriptor builder Adds a distinctly-named expert QMM implementation for the Gemma 4 26B-A4B MoE production shapes, gated by MLX_GATHER_QMM_EXPERT_SLICES: - qmm_t_expert_impl: BM32 expert tile body (BM16 fallback rows) taking a private/by-value row count; the shared qmm_t_impl constant-address ABI and all ordinary gathered/batched/dense QMM routes are unchanged. - build_gemma4_sorted_expert_tiles_bm32: one 128-thread threadgroup replaces the reference design's single-GPU-thread serial builder; parallel expert-range binary search, Hillis-Steele scan, and strided upper-bound descriptor emission. - Selector runs after the NAX-first route and requires affine BF16 transposed inputs, 4-bit gs=64 weights, 128 experts, assignment counts of exactly 4096/8192/16384, and the exact gate/up or down rank-3 shapes; every miss keeps the legacy route. NAX engagement is non-engagement, never bypassed. - device.{h,cpp}: one-shot request resolution, nonthrowing dual-symbol AOT probe/prewarm, relaxed-atomic diagnostics (requested, aotAvailable, naxAvailable, hits, per-class fallbacks). - gpu_tests: exact-shape arithmetic parity, fallback, and counter invariant probes. Retention standing (2026-08-09 production matrix): opt-in experiment. Standalone profile dropped (prefill -10.2% vs bracket); paired weighted-unsort+R1 profile retained-final (prefill +1.8%, TTFT -7.5%, decode +3.3%, arrival E2E +12.0%). NOTE: this source post-dates the benchmarked binaries/metallib (post-measurement kernel-body edit); rebuild and re-verify before any performance claim. * fix(mlx): fail-safe sortedness check in gemma expert tile builder; counter/atomic hygiene Review-wave fixes for the R1 expert-QMM path: - N1 (sortedness trust): build_gemma4_sorted_expert_tiles_bm32 now verifies each thread's post-binary-search segment boundary against the generalized invariant indices[start - 1] < lid <= indices[start] (edge threads check their single neighbor), votes per simdgroup via simd_or, folds the votes through threadgroup memory, and on any violation retracts count[0] to 0 (tile kernel then early-returns) and records the violation in count[1]; the buffer ABI is unchanged (count index 1 was previously unused). try_gemma4_expert_qmm allocates the second count element, drains the encoder after the builder, and re-routes a retracted call to the order-agnostic legacy path instead of dispatching the tile kernel (zero count is unambiguous: the selector's assignment gate guarantees M is 4096/8192/16384). - N2 (route-condition duplication): the sorted-RHS gate literal that appeared (negated) in the diagnostics record and in the dispatch decision is now the shared static constexpr predicate takes_sorted_rhs_route, so future tuning of the 16/4 thresholds cannot desynchronize counter vs route. - N3 (per-call bias normalization): gather_qmm_rhs no longer spends ensure_row_contiguous on biases before classification reads the raw tensor's fields; normalization runs only inside the winning-route branch (hit semantics unchanged; the legacy block keeps its own normalization point and ordering). - N4 (armed_ data race): Gemma4ExpertQMMCounters::armed_ is now std::atomic<bool> with relaxed loads/stores in armed(), snapshot(), snapshot_and_disarm() (read-then-write order preserved) and clear_and_arm(); the class remains non-copyable, now enforced. * fix(mlx): make the R1 sortedness fail-safe sound; proper retract attribution F1: the per-expert boundary vote was a partial detector -- an inversion inside a segment used by no other expert's boundary could escape, so "re-route on any violation" overclaimed. build_gemma4_sorted_expert_tiles_bm32 now also runs a strided adjacent-pair scan: thread lid checks indices[i-1] <= indices[i] for i = lid+1; i < M; i += 128, covering every adjacent pair in [1, M) exactly once (1..128 iterations at the reachable M in {4096,8192,16384}). Adjacent-pair monotonicity is transitive, so a clean scan is a sound and complete sortedness oracle; it folds into the same simd_or/threadgroup vote and the same retract (count[0]=0, count[1]=1). The boundary checks stay as cheap, precise diagnostics. F2: retracts were write-only in count[1] and surfaced as fallback_metallib_unavailable -- misattribution in the only observable surface. A dedicated fallback_sortedness_retracted counter now rides the GemmA4 route counters and the C diagnostics ABI (sizeof 80 -> 88, new uint64 at offset 80; existing offsets unchanged). try_gemma4_expert_qmm returns the route class: count[0]==0 with count[1]==1 records fallback_sortedness_retracted, any other unusable build keeps fallback_metallib_unavailable, then re-routes to the legacy path as before. F4: new doctest drives the full armed() -> clear_and_arm() -> snapshot_and_disarm() cycle and the attempts == hits + fallbacks invariant including the new class; the route-table and counter-invariant tests now cover fallback_sortedness_retracted. Verified: cmake tests 262/262 + 3550 assertions pass; metal -Wall -Wextra -fno-fast-math compile of kernels/quantized.metal is warning-free. * perf(metal): E=256 expert-tile route + trust + gpu::eval UAF fix — darkbloom-base mirror (#7) * perf(metal): instantiate E=256 expert-tile route for Qwen 3.5/3.6 MoE prefill (mirror of Cmlx/mlx 58fab46) * fix(metal): use-after-free in gpu::eval for primitives that synchronize mid-eval (mirror) * perf(metal): trust mode skips retract readback (mirror) * fix(compile): preserve all-cache binding cleanup --------- Co-authored-by: JasonHonKL <148705846+JasonHonKL@users.noreply.github.com> Co-authored-by: AK <144495202+AKnassa@users.noreply.github.com> Co-authored-by: Cheng <git@zcbenz.com> Co-authored-by: Adityaj0 <93090622+Adityaj0@users.noreply.github.com> Co-authored-by: anchor <codeanqiang@gmail.com> Co-authored-by: codeAnqiang-ma <273298913+codeAnqiang-ma@users.noreply.github.com> Co-authored-by: Rohan Gautam <rohan1gautam@gmail.com> Co-authored-by: Ayaan Gazali <ayaangazali.work@gmail.com> Co-authored-by: Erwin Zhang <59893706+erwinzhang7@users.noreply.github.com> Co-authored-by: katlun-lgtm <katlun@gmail.com> Co-authored-by: katlun-lgtm <264247399+katlun-lgtm@users.noreply.github.com> Co-authored-by: robertomeroni <150194833+robertomeroni@users.noreply.github.com> Co-authored-by: Fu Xiaonan <ht3fudatou@163.com> Co-authored-by: Fu Xiaonan <214359569+FU-max-boop@users.noreply.github.com> Co-authored-by: Hao Xu <hxu44@apple.com> Co-authored-by: Feli <89400571+FeliGame@users.noreply.github.com> Co-authored-by: Feli <feli@hnu.edu.cn> Co-authored-by: Eyüp Can Akman <eyupcanakman@gmail.com> Co-authored-by: Cheng <zcbenz@gmail.com> Co-authored-by: yentur <mr.yentur@gmail.com> Co-authored-by: Alessio Pollero <alessio.pollero@gmail.com> Co-authored-by: Zhiqi Zhang <zhiqizhangg@gmail.com> Co-authored-by: Daniel Hiltgen <dhiltgen@users.noreply.github.com> Co-authored-by: Tanish Jain <recklurker@gmail.com> Co-authored-by: hojin12312 <hojin12312@gmail.com> Co-authored-by: Xiang Chen <46052474+x14ngch3n@users.noreply.github.com> Co-authored-by: x14ngch3n <x14ngch3n@users.noreply.github.com> Co-authored-by: Duhyeon, Kim <49020301+dudududukim@users.noreply.github.com> Co-authored-by: rohith <kapellirohith@gmail.com> Co-authored-by: Ishaan Samantray <devteam.aegis@gmail.com> Co-authored-by: Aaishwarya Mishra <aaishwarymishra@gmail.com> Co-authored-by: Anastasiia Filippova <a_filippova@apple.com> Co-authored-by: Yanzhao Wang <19340816+wyanzhao@users.noreply.github.com> Co-authored-by: XXXXRT666 <157766680+XXXXRT666@users.noreply.github.com> Co-authored-by: Jake Bowhay <60778417+j-bowhay@users.noreply.github.com> Co-authored-by: vraj patel <87225460+vraj00222@users.noreply.github.com> Co-authored-by: Gusanidas <33495733+Gusanidas@users.noreply.github.com> Co-authored-by: Dwijen Patel <dwijen@gmail.com> Co-authored-by: Vladimir Iglovikov <ternaus@users.noreply.github.com> Co-authored-by: Brian C. <94733710+deBrian07@users.noreply.github.com> Co-authored-by: Daniel Hiltgen <daniel.hiltgen@ollama.com> Co-authored-by: YH Yan <strayberry0w0@gmail.com> Co-authored-by: katlun-lgtm <katlun@windyviews.com> Co-authored-by: anupsv <6407789+anupsv@users.noreply.github.com> Co-authored-by: Gajesh Naik <26431906+Gajesh2007@users.noreply.github.com> Co-authored-by: David Tai <davidtai@Davids-MBP.lan>
Addresses the conv2d manifestation of #3979 — the conv-side guard, not a fix for the general silent-error mechanism.
Problem
conv2don Metal silently returns all zeros for large batched inputs. The Winograd path keeps the padded input and both GEMM workspaces alive at once (~3.7x the input for the reported shapes), all referenced by one command buffer; once that working set exceedsrecommendedMaxWorkingSetSizethe driver fails the command buffer. As @jasonge27's instrumentation shows, Metal reports the OOM but errors from mid-eval commits are dropped, so nothing surfaces. Winograd's selection criteria (C % 32 == 0 && O % 32 == 0 && C + O >= 256) explain the channel dependence in the report. That dropped-error mechanism is general and needs a separate, complementary fix in the eval machinery; this PR avoids the OOM for conv2d, which also bounds its peak memory.Fix
Run the batch in tiles sized from the available working set, reusing one set of scratch buffers across tiles — mirroring the row tiling in
explicit_gemm_conv_ND_gpu. The budget is 3/4 ofrecommendedMaxWorkingSetSizeminus the conv's own tensors (input, output, and transformed filter), making the tile choice a deterministic function of the conv and the device. Measured failures start at ~99% of the reported limit; the quarter held back is headroom for what is not charged — GPU memory held by other processes and allocations unrelated to the conv.When the whole batch fits, the common case, the emitted work is unchanged. It falls back to the scratch-free implicit GEMM when not even one batch element fits, or when tiles would carry too few GEMM rows to amortize their fixed cost (small tiles can be far slower than the fallback).
Testing
The real threshold needs tens of GB, so tests shrink the budget instead:
MLX_CONV_WINOGRAD_WORKING_SETinjects the byte budget: covers automatic tiling and both fallbacks.MLX_CONV_WINOGRAD_TILE_BATCHforces a tile size (capped by the budget): covers uniform tiles, a short final tile, and a consumer op reading the tiled output in the same eval.Tiling keeps the per-element reduction order, so tiled Winograd is bit-identical to untiled while the implicit GEMM fallback never is — the tests use exact equality to pin each run to its intended path, with no memory measurements.
At real sizes on a 96 GB M2 Max: the issue's repro passes (previously all-zero at batch 80 on this machine; the reporter hit it at batch 40 on 48 GB), tiled vs untiled outputs are bit-identical at ~50 GB scale, CPU cross-checks agree, and the batch-80 repro peaks at 58 GB where untiled Winograd would need ~84 GB. On the reported fp32 shape (40x512x512, 192 -> 96 channels) Winograd runs ~1.7x faster than implicit GEMM (310 ms vs 536 ms), while a spatially tiny conv (256x4x4, 64 -> 192) cut into one-element Winograd tiles is ~40x slower than implicit GEMM — the two sides of the tile-versus-fallback trade-off. Mid-size Winograd shapes select a single tile and show no regression.