Make Scheduler::enqueue thread safe (#3423)

This commit is contained in:
Cheng
2026-04-20 14:30:05 +09:00
committed by GitHub
parent a6222f53d5
commit 1f5a413a27
3 changed files with 25 additions and 21 deletions
+17 -5
View File
@@ -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
View File
@@ -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
View File
@@ -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);
}