diff --git a/mlx/backend/cuda/allocator.cpp b/mlx/backend/cuda/allocator.cpp index 1cab6b4081..5f4b9f2d55 100644 --- a/mlx/backend/cuda/allocator.cpp +++ b/mlx/backend/cuda/allocator.cpp @@ -69,9 +69,9 @@ inline void* unified_malloc(size_t size) { inline void unified_free(void* data) { if (supports_managed_memory()) { - CHECK_CUDA_ERROR(cudaFree(data)); + cudaFree(data); } else { - CHECK_CUDA_ERROR(cudaFreeHost(data)); + cudaFreeHost(data); } } @@ -324,9 +324,9 @@ void CudaAllocator::free_async(CudaBuffer& buf, cudaStream_t stream) { if (!stream) { stream = free_streams_[buf.device]; } - CHECK_CUDA_ERROR(cudaFreeAsync(buf.data, stream)); + cudaFreeAsync(buf.data, stream); } else { - CHECK_CUDA_ERROR(cudaFree(buf.data)); + cudaFree(buf.data); } } } diff --git a/mlx/backend/cuda/cuda_utils.h b/mlx/backend/cuda/cuda_utils.h index f8a234ee65..083d19e6ff 100644 --- a/mlx/backend/cuda/cuda_utils.h +++ b/mlx/backend/cuda/cuda_utils.h @@ -26,11 +26,10 @@ class CudaHandle { } ~CudaHandle() { - // Skip if there was an error to avoid throwing in the destructors - if (cudaPeekAtLastError() != cudaSuccess) { - return; + // Skip error check to avoid throwing in destructors. + if (handle_ != nullptr) { + Destroy(handle_); } - reset(); } CudaHandle(const CudaHandle&) = delete; diff --git a/mlx/backend/cuda/device.cpp b/mlx/backend/cuda/device.cpp index e7d8f2620d..6bc16bf7ba 100644 --- a/mlx/backend/cuda/device.cpp +++ b/mlx/backend/cuda/device.cpp @@ -213,7 +213,11 @@ CommandEncoder::CommandEncoder(Device& d) } CommandEncoder::~CommandEncoder() { - synchronize(); + try { + synchronize(); + } catch (...) { + // Synchronizing can fail when the CUDA runtime is shutting down. + } worker_->stop(); } diff --git a/mlx/backend/cuda/event.cu b/mlx/backend/cuda/event.cu index d3b6f97f5d..bcea005443 100644 --- a/mlx/backend/cuda/event.cu +++ b/mlx/backend/cuda/event.cu @@ -232,8 +232,7 @@ AtomicEvent::AtomicEvent(Device& d) { cuda_free = cudaFree; coherent_ = false; } - buf_ = std::shared_ptr( - buf, [cuda_free](void* buf) { CHECK_CUDA_ERROR(cuda_free(buf)); }); + buf_ = std::shared_ptr(buf, [cuda_free](void* buf) { cuda_free(buf); }); if (coherent_) { *ptr() = 0; } else { diff --git a/mlx/compile.cpp b/mlx/compile.cpp index bb17c44962..3cab4c05d7 100644 --- a/mlx/compile.cpp +++ b/mlx/compile.cpp @@ -3,6 +3,7 @@ #include #include #include +#include #include #include #include @@ -302,9 +303,9 @@ std::uintptr_t get_function_address(const std::function& fun) { class CompileCache { public: struct CacheEntry { - CacheEntry(Stream stream, bool shapeless) - : stream(stream), shapeless(shapeless) {}; - Stream stream; + CacheEntry(std::optional stream, bool shapeless) + : stream(stream), shapeless(shapeless) {} + std::optional stream; bool shapeless; std::vector inputs; std::vector outputs; @@ -370,7 +371,7 @@ class CompileCache { // Loop over entries and check: // - Default stream and device match the entry's default stream // - Inputs match i.e. shapes and types must be equal. - auto stream = default_stream(default_device()); + auto stream = peek_default_stream(default_device()); for (CacheEntry& entry : entries) { // Check that the default stream and device match if (entry.stream != stream) { diff --git a/mlx/stream.cpp b/mlx/stream.cpp index b78ee67d67..902015cbf0 100644 --- a/mlx/stream.cpp +++ b/mlx/stream.cpp @@ -50,6 +50,10 @@ Stream default_stream(Device d) { return s.value(); } +std::optional peek_default_stream(Device d) { + return default_stream_storage(d); +} + void set_default_stream(Stream s) { if (!gpu::is_available() && s.device == Device::gpu) { throw std::invalid_argument( diff --git a/mlx/stream.h b/mlx/stream.h index 3099ee2682..331af60a61 100644 --- a/mlx/stream.h +++ b/mlx/stream.h @@ -2,6 +2,7 @@ #pragma once +#include #include #include @@ -29,6 +30,9 @@ struct MLX_API ThreadLocalStream : public Stream { /** Get the default stream of current thread for the given device. */ MLX_API Stream default_stream(Device d); +/** Get the default stream of current thread if it exists. */ +MLX_API std::optional peek_default_stream(Device d); + /** Make the stream the default for its device on current thread. */ MLX_API void set_default_stream(Stream s); diff --git a/mlx/transforms.cpp b/mlx/transforms.cpp index 1a8a5f9dc5..77d81e11e9 100644 --- a/mlx/transforms.cpp +++ b/mlx/transforms.cpp @@ -80,14 +80,16 @@ thread_local int detail::RetainGraph::tracing_counter{0}; array eval_impl(std::vector outputs, bool async) { std::deque tape; - // Make an effort to choose a good output stream - Stream stream = default_stream(default_device()); - for (auto& o : outputs) { - if (o.status() == array::Status::unscheduled && o.has_primitive()) { - stream = o.primitive().stream(); - break; + // Make an effort to choose a good output stream, and only create the default + // stream when there is no other choice. + Stream stream = [&outputs]() { + for (auto& o : outputs) { + if (o.status() == array::Status::unscheduled && o.has_primitive()) { + return o.primitive().stream(); + } } - } + return default_stream(default_device()); + }(); struct FenceInfo { int stream_index; diff --git a/tests/scheduler_tests.cpp b/tests/scheduler_tests.cpp index 8a98d35eb9..cc2228e0f3 100644 --- a/tests/scheduler_tests.cpp +++ b/tests/scheduler_tests.cpp @@ -114,6 +114,35 @@ TEST_CASE("test thread unsafe stream") { CHECK_EQ(expected, actual); } +TEST_CASE("test eval does not create default stream") { + auto s = new_thread_unsafe_stream(default_device()); + size_t num_streams = get_streams().size(); + + std::thread t([&] { + async_eval(arange(10, s)); + eval(arange(10, s)); + }); + t.join(); + + CHECK_EQ(get_streams().size(), num_streams); +} + +TEST_CASE("test compile does not create default stream") { + auto s = new_thread_unsafe_stream(default_device()); + size_t num_streams = get_streams().size(); + + std::function(const std::vector&)> fun = + [s](const std::vector& inputs) { + return std::vector{abs(inputs[0], s)}; + }; + auto cfun = compile(fun); + + std::thread t([&] { eval(cfun({array({-1, 2})})); }); + t.join(); + + CHECK_EQ(get_streams().size(), num_streams); +} + TEST_CASE("test thread local stream") { auto s = new_thread_local_stream(default_device()); int result = sum(arange(10, s)).item();