Skip to content
Merged
Show file tree
Hide file tree
Changes from 4 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
32 changes: 32 additions & 0 deletions benchmarks/python/synchronize_bench.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,32 @@
import time

import mlx.core as mx

world = mx.distributed.init()

a = mx.ones((5, 5), mx.int32)
its = 10
its_per_eval = 100


def fn(x):
for _ in range(its_per_eval):
x = mx.distributed.all_sum(x)
x = x - 1
return x


# warmup
for _ in range(5):
x = fn(a)
assert mx.array_equal(x, mx.ones_like(x))

tic = time.perf_counter()

for _ in range(its):
x = fn(a)
mx.eval(x)

toc = time.perf_counter()
ms = 1000 * (toc - tic) / (its * its_per_eval)
print(f"Time per iteration {ms:.6f} (ms)")
1 change: 1 addition & 0 deletions mlx/backend/metal/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,7 @@ target_sources(
${CMAKE_CURRENT_SOURCE_DIR}/distributed.cpp
${CMAKE_CURRENT_SOURCE_DIR}/device.cpp
${CMAKE_CURRENT_SOURCE_DIR}/event.cpp
${CMAKE_CURRENT_SOURCE_DIR}/fence.cpp
${CMAKE_CURRENT_SOURCE_DIR}/fft.cpp
${CMAKE_CURRENT_SOURCE_DIR}/hadamard.cpp
${CMAKE_CURRENT_SOURCE_DIR}/indexing.cpp
Expand Down
17 changes: 17 additions & 0 deletions mlx/backend/metal/device.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -134,6 +134,13 @@ CommandEncoder::~CommandEncoder() {
enc_->release();
}

void CommandEncoder::set_buffer(
const MTL::Buffer* buf,
int idx,
int64_t offset /* = 0 */) {
enc_->setBuffer(buf, offset, idx);
}

void CommandEncoder::set_input_array(
const array& a,
int idx,
Expand Down Expand Up @@ -189,6 +196,10 @@ void CommandEncoder::dispatch_threads(
enc_->dispatchThreads(grid_dims, group_dims);
}

void CommandEncoder::barrier() {
enc_->memoryBarrier(MTL::BarrierScopeBuffers);
}

Device::Device() {
auto pool = new_scoped_memory_pool();
device_ = load_device();
Expand Down Expand Up @@ -219,11 +230,16 @@ void Device::new_queue(int index) {
"[metal::Device] Failed to make new command queue.");
}
stream_map_.emplace(index, q);
stream_map_.at(index).event_fence = device_->newFence();
if (residency_set_ != nullptr) {
q->addResidencySet(residency_set_);
}
}

MTL::Fence* Device::get_event_fence(int index) {
return get_stream_(index).event_fence;
}

int Device::get_command_buffer_ops(int index) {
return get_stream_(index).buffer_ops;
}
Expand Down Expand Up @@ -338,6 +354,7 @@ CommandEncoder& Device::get_command_encoder(int index) {
if (stream.encoder == nullptr) {
stream.encoder = std::make_unique<CommandEncoder>(stream.buffer);
stream.fence = std::make_shared<Fence>(device_->newFence());
stream.encoder->wait_for_fence(stream.event_fence);

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I am kinda confused by this wait here. In the case where we don't use the Fence at all when would the event_fence be updated?

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I'm trying an alternative that doesn't require this fence cause I don't like it. But basically you can always wait on a fence. The wait will wait for any preceding calls to update. So if you never update it the wait it is essentially a no-op.

This fence is used to ensure no kernels start before the GPU signal kernel is done. So we update this fence when we signal from the GPU and then any command encoder that waits after that will wait for that update to finish.

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Ok, I pushed a change to get rid of this that should (and seems) to work, so you can disregard all the previous stuff with event_fence.

Basically it requires modifying the call to wait_gpu to take an array that we want to be sure is ready before any future kernels that depend on it run. It reuses our existing synchronization machinery (barriers + fences) and is nice in that it only encoders which actually depend on the output will wait for it.

I had to add a way to register_output_array since it's not actually part of a kernel.. but I think it's cleaner/ more efficient / and doesn't require this random stream_event which was very icky.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

So if you never update it the wait it is essentially a no-op.

What a horrible hidden API. I assumed wait without update is a deadlock hence the confusion before.

I had to add a way to register_output_array.

Yeah, so much better! Also only waiting on things that matter rather than everything on the whole stream. Plus avoid waiting on two fences which could have been the case before.

}
return *stream.encoder;
}
Expand Down
6 changes: 6 additions & 0 deletions mlx/backend/metal/device.h
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,7 @@ struct CommandEncoder {
void dispatch_threadgroups(MTL::Size grid_dims, MTL::Size group_dims);
void dispatch_threads(MTL::Size grid_dims, MTL::Size group_dims);
void maybeInsertBarrier();
void set_buffer(const MTL::Buffer* buf, int idx, int64_t offset = 0);

void set_compute_pipeline_state(MTL::ComputePipelineState* kernel) {
enc_->setComputePipelineState(kernel);
Expand Down Expand Up @@ -110,6 +111,8 @@ struct CommandEncoder {
return all_outputs_;
};

void barrier();

private:
MTL::ComputeCommandEncoder* enc_;
bool needs_barrier_{false};
Expand All @@ -136,6 +139,7 @@ struct DeviceStream {
if (buffer != nullptr) {
buffer->release();
}
event_fence->release();
};
MTL::CommandQueue* queue;
// A map of prior command encoder outputs to their corresponding fence
Expand All @@ -153,6 +157,7 @@ struct DeviceStream {
std::unique_ptr<CommandEncoder> encoder{nullptr};
std::shared_ptr<Fence> fence;
std::vector<array> temporaries;
MTL::Fence* event_fence;
};

class Device {
Expand All @@ -177,6 +182,7 @@ class Device {
void commit_command_buffer(int index);
CommandEncoder& get_command_encoder(int index);
void end_encoding(int index);
MTL::Fence* get_event_fence(int index);

void register_library(
const std::string& lib_name,
Expand Down
63 changes: 38 additions & 25 deletions mlx/backend/metal/distributed.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
#include "mlx/backend/common/utils.h"
#include "mlx/backend/metal/device.h"
#include "mlx/backend/metal/event.h"
#include "mlx/backend/metal/fence.h"
#include "mlx/distributed/ops.h"
#include "mlx/distributed/primitives.h"
#include "mlx/scheduler.h"
Expand All @@ -26,23 +27,28 @@ void AllReduce::eval_gpu(
assert(outputs.size() == 1);

auto& in = inputs[0];

Fence f{stream()};

if (in.event().valid()) {
f.update_gpu(in);
}
f.wait_gpu();

auto& out = outputs[0];
if (in.is_donatable()) {
out.move_shared_buffer(in);
} else {
out.set_data(allocator::malloc_or_wait(out.nbytes()));
}

auto e = Event(stream());
e.set_value(1);
signal_and_wait(in.event(), e);
auto task = [in = in,
out = out,
e = std::move(e),
f = std::move(f),
reduce_type = reduce_type_,
group = group()]() mutable {
if (in.event().valid()) {
in.event().wait();
f.wait();
}
switch (reduce_type) {
case Sum:
Expand All @@ -52,7 +58,7 @@ void AllReduce::eval_gpu(
default:
throw std::runtime_error("Only all reduce sum is supported for now");
}
e.signal();
f.update();
};
scheduler::enqueue(detail::communication_stream(), std::move(task));
}
Expand All @@ -67,17 +73,20 @@ void AllGather::eval_gpu(

out.set_data(allocator::malloc_or_wait(out.nbytes()));

auto e = Event(stream());
e.set_value(1);
signal_and_wait(in.event(), e);
Fence f{stream()};

if (in.event().valid()) {
f.update_gpu(in);
}
f.wait_gpu();

auto task =
[in = in, out = out, e = std::move(e), group = group()]() mutable {
[in = in, out = out, f = std::move(f), group = group()]() mutable {
if (in.event().valid()) {
in.event().wait();
f.wait();
}
distributed::detail::all_gather(group, in, out);
e.signal();
f.update();
};
scheduler::enqueue(detail::communication_stream(), std::move(task));
}
Expand All @@ -89,22 +98,28 @@ void Send::eval_gpu(
assert(outputs.size() == 1);

auto& in = inputs[0];

// Encode a signal event for the input
Fence f{stream()};
if (in.event().valid()) {
f.update_gpu(in);
}

auto& out = outputs[0];
move_or_copy(in, out);

// Schedule an async send on the comm stream
auto task = [in = in, out = out, group = group(), dst = dst_]() mutable {
auto task = [in = in,
out = out,
f = std::move(f),
group = group(),
dst = dst_]() mutable {
if (in.event().valid()) {
in.event().wait();
f.wait();
}
distributed::detail::send(group, out, dst);
};
scheduler::enqueue(detail::communication_stream(), std::move(task));

// Encode a signal event for the input
if (in.event().valid()) {
encode_signal(in.event());
}
}

void Recv::eval_gpu(
Expand All @@ -117,16 +132,14 @@ void Recv::eval_gpu(

out.set_data(allocator::malloc_or_wait(out.nbytes()));

auto e = Event(stream());
e.set_value(1);

encode_wait(e);
Fence f{stream()};
f.wait_gpu();

// Schedule an async recv on the comm stream
auto task =
[out = out, e = std::move(e), group = group(), src = src_]() mutable {
[out = out, f = std::move(f), group = group(), src = src_]() mutable {
distributed::detail::recv(group, out, src);
e.signal();
f.update();
};
scheduler::enqueue(detail::communication_stream(), std::move(task));
}
Expand Down
Loading