Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 4 additions & 4 deletions mlx/backend/cuda/allocator.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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);
}
}

Expand Down Expand Up @@ -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);
}
}
}
Expand Down
7 changes: 3 additions & 4 deletions mlx/backend/cuda/cuda_utils.h
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down
6 changes: 5 additions & 1 deletion mlx/backend/cuda/device.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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();
}

Expand Down
3 changes: 1 addition & 2 deletions mlx/backend/cuda/event.cu
Original file line number Diff line number Diff line change
Expand Up @@ -232,8 +232,7 @@ AtomicEvent::AtomicEvent(Device& d) {
cuda_free = cudaFree;
coherent_ = false;
}
buf_ = std::shared_ptr<void>(
buf, [cuda_free](void* buf) { CHECK_CUDA_ERROR(cuda_free(buf)); });
buf_ = std::shared_ptr<void>(buf, [cuda_free](void* buf) { cuda_free(buf); });
if (coherent_) {
*ptr() = 0;
} else {
Expand Down
9 changes: 5 additions & 4 deletions mlx/compile.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
#include <atomic>
#include <cstdlib>
#include <map>
#include <optional>
#include <shared_mutex>
#include <sstream>
#include <unordered_map>
Expand Down Expand Up @@ -302,9 +303,9 @@ std::uintptr_t get_function_address(const std::function<T(U...)>& fun) {
class CompileCache {
public:
struct CacheEntry {
CacheEntry(Stream stream, bool shapeless)
: stream(stream), shapeless(shapeless) {};
Stream stream;
CacheEntry(std::optional<Stream> stream, bool shapeless)
: stream(stream), shapeless(shapeless) {}
std::optional<Stream> stream;
bool shapeless;
std::vector<array> inputs;
std::vector<array> outputs;
Expand Down Expand Up @@ -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) {
Expand Down
4 changes: 4 additions & 0 deletions mlx/stream.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,10 @@ Stream default_stream(Device d) {
return s.value();
}

std::optional<Stream> 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(
Expand Down
4 changes: 4 additions & 0 deletions mlx/stream.h
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

#pragma once

#include <optional>
#include <tuple>
#include <vector>

Expand Down Expand Up @@ -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<Stream> peek_default_stream(Device d);

/** Make the stream the default for its device on current thread. */
MLX_API void set_default_stream(Stream s);

Expand Down
16 changes: 9 additions & 7 deletions mlx/transforms.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -80,14 +80,16 @@ thread_local int detail::RetainGraph::tracing_counter{0};
array eval_impl(std::vector<array> outputs, bool async) {
std::deque<array> 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;
Expand Down
29 changes: 29 additions & 0 deletions tests/scheduler_tests.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<std::vector<array>(const std::vector<array>&)> fun =
[s](const std::vector<array>& inputs) {
return std::vector<array>{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<int>();
Expand Down
Loading