From 9fba0906eee536ec4ef358b93f128dbb4ed953e5 Mon Sep 17 00:00:00 2001 From: Tarjei Mandt Date: Fri, 7 Aug 2026 14:56:02 +1000 Subject: [PATCH] Remove the per-op completion handler in gpu::eval --- mlx/backend/metal/device.cpp | 3 +++ mlx/backend/metal/device.h | 10 ++++++++++ mlx/backend/metal/eval.cpp | 26 ++++++++++---------------- 3 files changed, 23 insertions(+), 16 deletions(-) diff --git a/mlx/backend/metal/device.cpp b/mlx/backend/metal/device.cpp index 197b919b1d..c7f0eaed24 100644 --- a/mlx/backend/metal/device.cpp +++ b/mlx/backend/metal/device.cpp @@ -520,6 +520,7 @@ void CommandEncoder::commit(std::function completion) { [&error_ = error_, wait_events = std::move(wait_events_), signal_events = std::move(signal_events_), + retained = std::move(retained_buffers_), completion = std::move(completion)](MTL::CommandBuffer* cbuf) { if (completion) { completion(); @@ -556,6 +557,8 @@ void CommandEncoder::commit(std::function completion) { buffer_ = NS::RetainPtr(queue_->commandBufferWithUnretainedReferences()); buffer_ops_ = 0; buffer_sizes_ = 0; + retained_buffers_.clear(); + retained_ptrs_.clear(); } void CommandEncoder::synchronize() { diff --git a/mlx/backend/metal/device.h b/mlx/backend/metal/device.h index 871e95ccff..4ba79cebb6 100644 --- a/mlx/backend/metal/device.h +++ b/mlx/backend/metal/device.h @@ -54,6 +54,14 @@ class MLX_API CommandEncoder { void add_temporary(array arr); void add_temporaries(std::vector arrays); + // Keep a buffer alive until the current command buffer completes. Deduped + // because a primitive may bind the same array in several slots. + void hold_buffer(const std::shared_ptr& buf) { + if (buf && retained_ptrs_.insert(buf.get()).second) { + retained_buffers_.push_back(buf); + } + } + void dispatch_threadgroups(MTL::Size grid_dims, MTL::Size group_dims); void dispatch_threads(MTL::Size grid_dims, MTL::Size group_dims); void maybeInsertBarrier(); @@ -129,6 +137,8 @@ class MLX_API CommandEncoder { bool needs_barrier_{false}; bool concurrent_{false}; std::vector temporaries_; + std::vector> retained_buffers_; + std::unordered_set retained_ptrs_; std::unordered_set prev_inputs_; std::unordered_set prev_outputs_; std::unordered_set next_inputs_; diff --git a/mlx/backend/metal/eval.cpp b/mlx/backend/metal/eval.cpp index e6826253de..36c120bf70 100644 --- a/mlx/backend/metal/eval.cpp +++ b/mlx/backend/metal/eval.cpp @@ -1,6 +1,4 @@ // Copyright © 2023-2024 Apple Inc. -#include - #include "mlx/backend/gpu/eval.h" #include "mlx/backend/metal/device.h" #include "mlx/backend/metal/utils.h" @@ -44,27 +42,23 @@ void eval(array& arr) { debug_set_primitive_buffer_label(command_buffer, arr.primitive()); arr.primitive().eval_gpu(arr.inputs(), outputs); } - std::unordered_set> buffers; + // Skip a donated output's buffer since holding it blocks allocator reuse. + const auto& out_data = arr.data_shared_ptr(); for (auto& in : arr.inputs()) { - buffers.insert(in.data_shared_ptr()); - } - for (auto& s : arr.siblings()) { - buffers.insert(s.data_shared_ptr()); + if (in.data_shared_ptr() != out_data) { + encoder.hold_buffer(in.data_shared_ptr()); + } } - // Remove the output if it was donated to by an input - if (auto it = buffers.find(arr.data_shared_ptr()); it != buffers.end()) { - buffers.erase(it); + for (auto& sib : arr.siblings()) { + if (sib.data_shared_ptr() != out_data) { + encoder.hold_buffer(sib.data_shared_ptr()); + } } if (encoder.needs_commit()) { encoder.end_encoding(); scheduler::notify_new_task(s); - encoder.commit([s, buffers = std::move(buffers)]() { - scheduler::notify_task_completion(s); - }); - } else { - command_buffer->addCompletedHandler( - [buffers = std::move(buffers)](MTL::CommandBuffer* cbuf) {}); + encoder.commit([s]() { scheduler::notify_task_completion(s); }); } }