diff --git a/mlx/backend/common/error.h b/mlx/backend/common/error.h new file mode 100644 index 0000000000..ba1164f192 --- /dev/null +++ b/mlx/backend/common/error.h @@ -0,0 +1,49 @@ +// Copyright © 2026 Apple Inc. + +#pragma once + +#include +#include +#include + +namespace mlx::core { + +class Error { + public: + // TODO: Use std::atomic when it gets supported in Xcode. + using Message = std::shared_ptr; + + void set_message(Message msg) { + std::atomic_store(&message_, std::move(msg)); + } + + bool valid() const { + auto msg = std::atomic_load(&message_); + return msg.get(); + } + + // If |ptr| is a valid event, copy and return true. + bool store_if_valid(const Error* ptr) { + if (ptr && this != ptr) { + Message msg = std::atomic_load(&ptr->message_); + if (msg) { + set_message(std::move(msg)); + return true; + } + } + return false; + } + + // If current error is valid, throw and clear. + void check() { + auto msg = std::atomic_exchange(&message_, {}); + if (msg) { + throw std::runtime_error(*msg); + } + } + + private: + Message message_; +}; + +} // namespace mlx::core diff --git a/mlx/backend/metal/device.cpp b/mlx/backend/metal/device.cpp index cd601dd319..dfee91e767 100644 --- a/mlx/backend/metal/device.cpp +++ b/mlx/backend/metal/device.cpp @@ -529,20 +529,22 @@ void CommandEncoder::commit(std::function completion) { } // If any of the waited event has error in it, poison the encoder. for (auto& event : wait_events) { - if (event->error()) { - error_ = event->error(); + if (error_.store_if_valid(event->error())) { break; } } // Set error only when no error happended before, to preserve the // earliest error. - if (!error_ && cbuf->status() == MTL::CommandBufferStatusError) { - error_ = std::make_shared(fmt::format( - "[METAL] Command buffer execution failed: {}.", - cbuf->error()->localizedDescription()->utf8String())); + bool has_error = error_.valid(); + if (!has_error && cbuf->status() == MTL::CommandBufferStatusError) { + error_.set_message( + std::make_shared(fmt::format( + "[METAL] Command buffer execution failed: {}.", + cbuf->error()->localizedDescription()->utf8String()))); + has_error = true; } // Poison all the signaled events when error happened. - if (error_) { + if (has_error) { for (auto& [event, value] : signal_events) { event->set_error(error_); } @@ -568,20 +570,17 @@ void CommandEncoder::synchronize() { commit(); cbuf->waitUntilCompleted(); - if (error_ && !exiting_) { - auto error = std::move(error_); - throw std::runtime_error(*error); + if (!exiting_) { + error_.check(); } } MTL::ComputeCommandEncoder* CommandEncoder::get_command_encoder() { if (!encoder_) { + error_.check(); encoder_ = NS::RetainPtr( buffer_->computeCommandEncoder(MTL::DispatchTypeConcurrent)); fence_ = NS::TransferPtr(device_.mtl_device()->newFence()); - // Reset error when user starts to encode new commands, they are supposed to - // have handled the error in synchronize() or Event::wait(). - error_.reset(); } return encoder_.get(); } diff --git a/mlx/backend/metal/device.h b/mlx/backend/metal/device.h index 871e95ccff..6c642b4b03 100644 --- a/mlx/backend/metal/device.h +++ b/mlx/backend/metal/device.h @@ -6,11 +6,11 @@ #include #include #include -#include #include #include #include "mlx/array.h" +#include "mlx/backend/common/error.h" #include "mlx/backend/common/metal_kernel.h" #include "mlx/backend/metal/resident.h" #include "mlx/device.h" @@ -119,7 +119,7 @@ class MLX_API CommandEncoder { std::vector, uint64_t>> signal_events_; // Error from previous commited command buffer. - std::shared_ptr error_; + Error error_; // Encoder for issuing GPU commands. // The members are used within a single ComputeCommandEncoder and will be diff --git a/mlx/backend/metal/event.cpp b/mlx/backend/metal/event.cpp index 77f48f0838..0529a72514 100644 --- a/mlx/backend/metal/event.cpp +++ b/mlx/backend/metal/event.cpp @@ -26,24 +26,21 @@ EventImpl::~EventImpl() { } void EventImpl::wait(uint64_t value) { - check_error(); + if (auto* p = error(); p) { + p->check(); + } mtl_event_->waitUntilSignaledValue(value, -1); // never times out - check_error(); + if (auto* p = error(); p) { + p->check(); + } } void EventImpl::signal(uint64_t value) { mtl_event_->setSignaledValue(value); } -void EventImpl::set_error(std::shared_ptr error) { - std::atomic_store(&error_, std::move(error)); -} - -void EventImpl::check_error() { - auto error = std::atomic_exchange(&error_, {}); - if (error) { - throw std::runtime_error(*error); - } +void EventImpl::set_error(Error& error) { + error_.store(&error); } } // namespace metal diff --git a/mlx/backend/metal/event.h b/mlx/backend/metal/event.h index c5c82a7cd3..ea0ebdcac0 100644 --- a/mlx/backend/metal/event.h +++ b/mlx/backend/metal/event.h @@ -12,11 +12,10 @@ class EventImpl { void wait(uint64_t value); void signal(uint64_t value); - void set_error(std::shared_ptr error); - void check_error(); + void set_error(Error& error); - const auto& error() const { - return error_; + Error* error() const { + return error_.load(); } auto* mtl_event() { @@ -24,8 +23,8 @@ class EventImpl { } private: - // TODO: Use std::atomic when it gets supported in Xcode. - std::shared_ptr error_; + // All streams outlive events so pointers would be always valid. + std::atomic error_; NS::SharedPtr mtl_event_; };