#include "triton/common/thread_pool.h"
#include <stdexcept>
namespace triton { namespace common {
ThreadPool::ThreadPool(size_t thread_count)
{
if (!thread_count) {
throw std::invalid_argument("Thread count must be greater than zero.");
}
const auto worker_loop = [this]() {
while (true) {
Task task;
{
std::unique_lock<std::mutex> lk(queue_mtx_);
cv_.wait(lk, [&]() { return !task_queue_.empty() || stop_; });
if (stop_ && task_queue_.empty()) {
break;
}
task = std::move(task_queue_.front());
task_queue_.pop();
}
if (task) {
task();
}
}
};
workers_.reserve(thread_count);
for (size_t i = 0; i < thread_count; ++i) {
workers_.emplace_back(worker_loop);
}
}
ThreadPool::~ThreadPool()
{
{
std::lock_guard<std::mutex> lk(queue_mtx_);
stop_ = true;
}
cv_.notify_all();
for (auto& t : workers_) {
t.join();
}
}
void
ThreadPool::enqueue(Task&& task)
{
{
std::lock_guard<std::mutex> lk(queue_mtx_);
if (stop_) {
return;
}
task_queue_.push(std::move(task));
}
cv_.notify_one();
}
}}