From e5f82010f34b558107b037845f0bfb77809c43ba Mon Sep 17 00:00:00 2001 From: Anastasiia Filippova Date: Wed, 23 Sep 2026 01:02:12 +0200 Subject: [PATCH 1/4] fist attempt: Fence::update() marks the input --- mlx/backend/cuda/fence.cpp | 31 ++++++++++++++++--------------- mlx/backend/metal/fence.cpp | 19 +++++++++++-------- mlx/backend/no_gpu/fence.cpp | 9 ++++++--- mlx/transforms.cpp | 18 ++++++++++++------ 4 files changed, 45 insertions(+), 32 deletions(-) diff --git a/mlx/backend/cuda/fence.cpp b/mlx/backend/cuda/fence.cpp index 39c2de6891..097ffeae7e 100644 --- a/mlx/backend/cuda/fence.cpp +++ b/mlx/backend/cuda/fence.cpp @@ -5,19 +5,16 @@ #include "mlx/backend/cuda/device.h" #include "mlx/backend/cuda/event.h" +#include + namespace mlx::core { struct FenceImpl { uint32_t count; - Event gpu_event; + std::vector gpu_events; Event cpu_event; FenceImpl(uint32_t count, Stream s) : count(count), cpu_event(s) { - if (s.device == Device::gpu) { - gpu_event = Event(s); - // A value of one selects a native CUDA event. - gpu_event.set_value(1); - } // Ensure that we use AtomicEvent, it is the only event that can order a CPU // stream against the GPU. cpu_event.cast().ensure_created(s, 2); @@ -28,21 +25,23 @@ Fence::Fence(Stream s) { fence_ = std::make_shared(0, s); } -void Fence::wait(Stream s, const array&) { +void Fence::wait(Stream s, const array&, uint32_t value) { auto& f = cast(); - if (f.count == 0) { + if (value == 0) { return; } - if (f.gpu_event.valid() && s.device == Device::gpu) { - f.gpu_event.wait(s); + if (!f.gpu_events.empty() && s.device == Device::gpu) { + f.gpu_events.at(value - 1).wait(s); } else { // AtomicEvent can not reliably notify a GPU stream, so a dependency that // involves the CPU keeps the synchronous wait. - f.cpu_event.wait(); + auto event = f.cpu_event; + event.set_value(value); + event.wait(); } } -void Fence::update(Stream s, const array& a, bool cross_device) { +uint32_t Fence::update(Stream s, const array& a, bool cross_device) { auto& f = cast(); if (cross_device) { // Move to managed memory if there is a device switch @@ -56,12 +55,14 @@ void Fence::update(Stream s, const array& a, bool cross_device) { } f.count++; if (s.device == Device::gpu) { - f.gpu_event.signal(s); + // Keep each recording so a consumer can wait for an earlier update. + auto& event = f.gpu_events.emplace_back(s); + event.set_value(1); + event.signal(s); } - // The counted event stays current, so a CPU consumer is always ordered - // against every update. f.cpu_event.set_value(f.count); f.cpu_event.signal(s); + return f.count; } } // namespace mlx::core diff --git a/mlx/backend/metal/fence.cpp b/mlx/backend/metal/fence.cpp index 70dd0e33bd..41b8465068 100644 --- a/mlx/backend/metal/fence.cpp +++ b/mlx/backend/metal/fence.cpp @@ -44,18 +44,20 @@ Fence::Fence(Stream stream) { fence_ = std::make_shared(stream); } -void Fence::wait(Stream stream, const array& x) { +void Fence::wait(Stream stream, const array& x, uint32_t value) { auto& f = *static_cast(fence_.get()); if (!f.use_fast) { - f.event->wait(stream); + auto event = *f.event; + event.set_value(value); + event.wait(stream); return; } if (stream.device == Device::cpu) { - scheduler::enqueue(stream, [fence_ = fence_, count = f.count]() mutable { + scheduler::enqueue(stream, [fence_ = fence_, value]() mutable { auto& f = *static_cast(fence_.get()); - while (f.cpu_value()[0] < count) { + while (f.cpu_value()[0] < value) { } }); return; @@ -74,21 +76,21 @@ void Fence::wait(Stream stream, const array& x) { auto buf = static_cast(f.fence); compute_encoder.set_buffer(buf, 0); - compute_encoder.set_bytes(f.count, 1); + compute_encoder.set_bytes(value, 1); compute_encoder.dispatch_threads(kernel_dims, kernel_dims); compute_encoder.get_command_buffer()->addCompletedHandler( [fence_ = fence_](MTL::CommandBuffer* cbuf) {}); } -void Fence::update(Stream stream, const array& x, bool cross_device) { +uint32_t Fence::update(Stream stream, const array& x, bool cross_device) { auto& f = *static_cast(fence_.get()); f.count++; if (!f.use_fast) { f.event->set_value(f.count); f.event->signal(stream); - return; + return f.count; } if (stream.device == Device::cpu) { @@ -96,7 +98,7 @@ void Fence::update(Stream stream, const array& x, bool cross_device) { auto& f = *static_cast(fence_.get()); f.cpu_value()[0] = count; }); - return; + return f.count; } auto& d = metal::device(stream.device); @@ -130,6 +132,7 @@ void Fence::update(Stream stream, const array& x, bool cross_device) { compute_encoder.get_command_buffer()->addCompletedHandler( [fence_ = fence_](MTL::CommandBuffer* cbuf) {}); + return f.count; } } // namespace mlx::core diff --git a/mlx/backend/no_gpu/fence.cpp b/mlx/backend/no_gpu/fence.cpp index 05852c860b..2e027032d9 100644 --- a/mlx/backend/no_gpu/fence.cpp +++ b/mlx/backend/no_gpu/fence.cpp @@ -16,15 +16,18 @@ Fence::Fence(Stream s) { fence_ = std::make_shared(0, s); } -void Fence::wait(Stream s, const array&) { - cast().event.wait(s); +void Fence::wait(Stream s, const array&, uint32_t value) { + auto event = cast().event; + event.set_value(value); + event.wait(s); } -void Fence::update(Stream s, const array&, bool) { +uint32_t Fence::update(Stream s, const array&, bool) { auto& f = cast(); f.count++; f.event.set_value(f.count); f.event.signal(s); + return f.count; } } // namespace mlx::core diff --git a/mlx/transforms.cpp b/mlx/transforms.cpp index 694bca53d8..1a8a5f9dc5 100644 --- a/mlx/transforms.cpp +++ b/mlx/transforms.cpp @@ -89,8 +89,12 @@ array eval_impl(std::vector outputs, bool async) { } } - // Map of array id that needs fence and stream it's computed on - std::unordered_map> needs_fence; + struct FenceInfo { + int stream_index; + bool cross_device; + uint32_t value{0}; + }; + std::unordered_map needs_fence; auto synchronizer = array( {}, bool_, std::make_shared(stream), std::move(outputs)); @@ -146,9 +150,9 @@ array eval_impl(std::vector outputs, bool async) { a.primitive().stream().device != in.primitive().stream().device; auto [it, inserted] = needs_fence.emplace( in.id(), - std::make_pair(in.primitive().stream().index, device_switch)); + FenceInfo{in.primitive().stream().index, device_switch}); if (!inserted) { - it->second.second |= device_switch; + it->second.cross_device |= device_switch; } } } @@ -251,7 +255,8 @@ array eval_impl(std::vector outputs, bool async) { // Use fence to wait within a single eval // Get the input array's stream fence and wait on the // output arrays stream - fences[it->second.first].wait(stream, in); + auto& info = it->second; + fences.at(info.stream_index).wait(stream, in, info.value); } else if (in.event().valid()) { if (in.event().is_signaled()) { in.detach_event(); @@ -291,7 +296,8 @@ array eval_impl(std::vector outputs, bool async) { if (it == fences.end()) { it = fences.emplace(stream.index, Fence{stream}).first; } - it->second.update(stream, a, nf->second.second); + nf->second.value = + it->second.update(stream, a, nf->second.cross_device); } }; From b8866f9cf87f1883473ffd1f254814a1ae37cf81 Mon Sep 17 00:00:00 2001 From: Anastasiia Filippova Date: Wed, 23 Sep 2026 01:30:48 +0200 Subject: [PATCH 2/4] fix declaration --- mlx/fence.h | 15 ++++++--------- 1 file changed, 6 insertions(+), 9 deletions(-) diff --git a/mlx/fence.h b/mlx/fence.h index 3fd5da333b..fd6bc116a4 100644 --- a/mlx/fence.h +++ b/mlx/fence.h @@ -1,6 +1,6 @@ // Copyright © 2024 Apple Inc. -#include +#include #include "mlx/array.h" @@ -8,12 +8,9 @@ namespace mlx::core { /* A fence to be used for synchronizing work between streams. * - * Calls to `wait` wait in the given stream until all previous calls to update - * are complete on their given stream. - * - * The array passed to `update` is computed and visible after the call to - * `wait` returns. The array passed to `wait` will not be read until all - * previous calls to `update` have completed. + * `update` returns a value that marks when its array is computed and visible. + * `wait` orders work in the consumer stream after that value is signaled. + * Later updates do not extend an earlier array's wait. * * Note, calls to `update` should always be from the same thread or explicitly * synchronized so that they occur in sequence. Calls to `wait` can be on any @@ -29,8 +26,8 @@ class Fence { Fence() {}; explicit Fence(Stream stream); - void update(Stream stream, const array& x, bool cross_device); - void wait(Stream stream, const array& x); + uint32_t update(Stream stream, const array& x, bool cross_device); + void wait(Stream stream, const array& x, uint32_t value); template auto& cast() const { From 125f7cb1f81289b7e419ec27fe93ea8b02a495dd Mon Sep 17 00:00:00 2001 From: Anastasiia Filippova Date: Thu, 24 Sep 2026 16:20:35 +0200 Subject: [PATCH 3/4] Update mlx/backend/metal/fence.cpp Co-authored-by: Cheng --- mlx/backend/metal/fence.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mlx/backend/metal/fence.cpp b/mlx/backend/metal/fence.cpp index 41b8465068..6594f55070 100644 --- a/mlx/backend/metal/fence.cpp +++ b/mlx/backend/metal/fence.cpp @@ -48,7 +48,7 @@ void Fence::wait(Stream stream, const array& x, uint32_t value) { auto& f = *static_cast(fence_.get()); if (!f.use_fast) { - auto event = *f.event; + auto& event = *f.event; event.set_value(value); event.wait(stream); return; From b0aaa906c305a7c60359da05dc4b5be421390c1a Mon Sep 17 00:00:00 2001 From: Anastasiia Filippova Date: Thu, 24 Sep 2026 16:20:47 +0200 Subject: [PATCH 4/4] Update mlx/backend/cuda/fence.cpp Co-authored-by: Cheng --- mlx/backend/cuda/fence.cpp | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mlx/backend/cuda/fence.cpp b/mlx/backend/cuda/fence.cpp index 097ffeae7e..e775d73ece 100644 --- a/mlx/backend/cuda/fence.cpp +++ b/mlx/backend/cuda/fence.cpp @@ -35,7 +35,7 @@ void Fence::wait(Stream s, const array&, uint32_t value) { } else { // AtomicEvent can not reliably notify a GPU stream, so a dependency that // involves the CPU keeps the synchronous wait. - auto event = f.cpu_event; + auto& event = f.cpu_event; event.set_value(value); event.wait(); }