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
31 changes: 16 additions & 15 deletions mlx/backend/cuda/fence.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -5,19 +5,16 @@
#include "mlx/backend/cuda/device.h"
#include "mlx/backend/cuda/event.h"

#include <vector>

namespace mlx::core {

struct FenceImpl {
uint32_t count;
Event gpu_event;
std::vector<Event> 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<cu::EventImpl>().ensure_created(s, 2);
Expand All @@ -28,21 +25,23 @@ Fence::Fence(Stream s) {
fence_ = std::make_shared<FenceImpl>(0, s);
}

void Fence::wait(Stream s, const array&) {
void Fence::wait(Stream s, const array&, uint32_t value) {
auto& f = cast<FenceImpl>();
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<FenceImpl>();
if (cross_device) {
// Move to managed memory if there is a device switch
Expand All @@ -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
19 changes: 11 additions & 8 deletions mlx/backend/metal/fence.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -44,18 +44,20 @@ Fence::Fence(Stream stream) {
fence_ = std::make_shared<FenceImpl>(stream);
}

void Fence::wait(Stream stream, const array& x) {
void Fence::wait(Stream stream, const array& x, uint32_t value) {
auto& f = *static_cast<FenceImpl*>(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<FenceImpl*>(fence_.get());
while (f.cpu_value()[0] < count) {
while (f.cpu_value()[0] < value) {
}
});
return;
Expand All @@ -74,29 +76,29 @@ void Fence::wait(Stream stream, const array& x) {

auto buf = static_cast<MTL::Buffer*>(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<FenceImpl*>(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) {
scheduler::enqueue(stream, [fence_ = fence_, count = f.count]() mutable {
auto& f = *static_cast<FenceImpl*>(fence_.get());
f.cpu_value()[0] = count;
});
return;
return f.count;
}

auto& d = metal::device(stream.device);
Expand Down Expand Up @@ -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
9 changes: 6 additions & 3 deletions mlx/backend/no_gpu/fence.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -16,15 +16,18 @@ Fence::Fence(Stream s) {
fence_ = std::make_shared<FenceImpl>(0, s);
}

void Fence::wait(Stream s, const array&) {
cast<FenceImpl>().event.wait(s);
void Fence::wait(Stream s, const array&, uint32_t value) {
auto event = cast<FenceImpl>().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<FenceImpl>();
f.count++;
f.event.set_value(f.count);
f.event.signal(s);
return f.count;
}

} // namespace mlx::core
15 changes: 6 additions & 9 deletions mlx/fence.h
Original file line number Diff line number Diff line change
@@ -1,19 +1,16 @@
// Copyright © 2024 Apple Inc.

#include <vector>
#include <cstdint>

#include "mlx/array.h"

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
Expand All @@ -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 <typename T>
auto& cast() const {
Expand Down
18 changes: 12 additions & 6 deletions mlx/transforms.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -89,8 +89,12 @@ array eval_impl(std::vector<array> outputs, bool async) {
}
}

// Map of array id that needs fence and stream it's computed on
std::unordered_map<uintptr_t, std::pair<uint32_t, bool>> needs_fence;
struct FenceInfo {
int stream_index;
bool cross_device;
uint32_t value{0};
};
std::unordered_map<uintptr_t, FenceInfo> needs_fence;

auto synchronizer = array(
{}, bool_, std::make_shared<Synchronizer>(stream), std::move(outputs));
Expand Down Expand Up @@ -146,9 +150,9 @@ array eval_impl(std::vector<array> 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;
}
}
}
Expand Down Expand Up @@ -251,7 +255,8 @@ array eval_impl(std::vector<array> 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();
Expand Down Expand Up @@ -291,7 +296,8 @@ array eval_impl(std::vector<array> 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);
}
};

Expand Down
Loading