Make Scheduler::enqueue thread safe (#3423)
This commit is contained in:
+17
-5
@@ -32,12 +32,24 @@ Scheduler::Scheduler() {
|
||||
|
||||
Scheduler::~Scheduler() = default;
|
||||
|
||||
void Scheduler::new_thread(Device::DeviceType type) {
|
||||
if (type == Device::gpu) {
|
||||
threads_.push_back(nullptr);
|
||||
} else {
|
||||
threads_.push_back(std::make_unique<StreamThread>());
|
||||
void Scheduler::enqueue(Stream s, std::function<void()> task) {
|
||||
StreamThread* st = nullptr;
|
||||
{
|
||||
std::shared_lock lock(threads_mtx_);
|
||||
auto it = threads_.find(s.index);
|
||||
if (it != threads_.end()) {
|
||||
st = it->second.get();
|
||||
}
|
||||
}
|
||||
if (!st) {
|
||||
std::unique_lock lock(threads_mtx_);
|
||||
auto it = threads_.find(s.index);
|
||||
if (it == threads_.end()) {
|
||||
it = threads_.emplace(s.index, std::make_unique<StreamThread>()).first;
|
||||
}
|
||||
st = it->second.get();
|
||||
}
|
||||
st->enqueue(std::move(task));
|
||||
}
|
||||
|
||||
/** A singleton scheduler to manage devices, streams, and task execution. */
|
||||
|
||||
+7
-14
@@ -5,6 +5,7 @@
|
||||
#include <atomic>
|
||||
#include <future>
|
||||
#include <queue>
|
||||
#include <shared_mutex>
|
||||
#include <thread>
|
||||
#include <unordered_map>
|
||||
|
||||
@@ -50,21 +51,20 @@ struct StreamThread {
|
||||
}
|
||||
}
|
||||
|
||||
template <typename F>
|
||||
void enqueue(F&& f) {
|
||||
void enqueue(std::function<void()> f) {
|
||||
{
|
||||
std::lock_guard<std::mutex> lk(mtx);
|
||||
if (stop) {
|
||||
throw std::runtime_error(
|
||||
"Cannot enqueue work after stream is stopped.");
|
||||
}
|
||||
q.emplace(std::forward<F>(f));
|
||||
q.emplace(std::move(f));
|
||||
}
|
||||
cond.notify_one();
|
||||
}
|
||||
};
|
||||
|
||||
class Scheduler {
|
||||
class MLX_API Scheduler {
|
||||
public:
|
||||
Scheduler();
|
||||
~Scheduler();
|
||||
@@ -75,8 +75,7 @@ class Scheduler {
|
||||
Scheduler& operator=(const Scheduler&) = delete;
|
||||
Scheduler& operator=(Scheduler&&) = delete;
|
||||
|
||||
template <typename F>
|
||||
void enqueue(const Stream& stream, F&& f);
|
||||
void enqueue(Stream s, std::function<void()> task);
|
||||
|
||||
void notify_new_task(const Stream& stream) {
|
||||
{
|
||||
@@ -111,19 +110,13 @@ class Scheduler {
|
||||
private:
|
||||
friend Stream mlx::core::new_stream(Device d);
|
||||
|
||||
void new_thread(Device::DeviceType type);
|
||||
|
||||
int n_active_tasks_{0};
|
||||
std::vector<std::unique_ptr<StreamThread>> threads_;
|
||||
std::unordered_map<int, std::unique_ptr<StreamThread>> threads_;
|
||||
std::shared_mutex threads_mtx_;
|
||||
std::condition_variable completion_cv;
|
||||
std::mutex mtx;
|
||||
};
|
||||
|
||||
template <typename F>
|
||||
void Scheduler::enqueue(const Stream& stream, F&& f) {
|
||||
threads_[stream.index]->enqueue(std::forward<F>(f));
|
||||
}
|
||||
|
||||
MLX_API Scheduler& scheduler();
|
||||
|
||||
template <typename F>
|
||||
|
||||
+1
-2
@@ -3,7 +3,7 @@
|
||||
#include "mlx/stream.h"
|
||||
#include "mlx/backend/cpu/device_info.h"
|
||||
#include "mlx/backend/gpu/device_info.h"
|
||||
#include "mlx/scheduler.h"
|
||||
#include "mlx/backend/gpu/eval.h"
|
||||
|
||||
#include <array>
|
||||
#include <map>
|
||||
@@ -68,7 +68,6 @@ Stream new_stream(Device d) {
|
||||
std::unique_lock lock(mtx);
|
||||
int index = streams.size();
|
||||
auto& s = streams.emplace_back(index, d);
|
||||
scheduler::scheduler().new_thread(d.type);
|
||||
if (d == Device::gpu) {
|
||||
gpu::new_stream(s);
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user