LIMA
Libre Multilingual Analyzer — C++ API
Loading...
Searching...
No Matches
thread_pool.h
Go to the documentation of this file.
1// Copyright 2021 CEA LIST
2// SPDX-FileCopyrightText: 2022 CEA LIST <gael.de-chalendar@cea.fr>
3//
4// SPDX-License-Identifier: MIT
5
6#ifndef DEEPLIMA_SRC_INCLUDE_THREAD_POOL_H
7#define DEEPLIMA_SRC_INCLUDE_THREAD_POOL_H
8
9#include <iostream>
10#include <assert.h>
11#include <queue>
12#include <functional>
13#include <thread>
14#include <atomic>
15#include <mutex>
16#include <condition_variable>
17
18namespace deeplima
19{
20
21template <class P>
23{
24public:
25
26 ThreadPool(size_t num_threads = 0)
27 : m_stop(false)
28 {
29 if (num_threads > 0)
30 {
31 init(num_threads);
32 }
33 }
34
35 void init(size_t num_threads)
36 {
37 assert(num_threads > 0);
38 assert(m_workers.size() == 0);
39
40 m_workers.reserve(num_threads);
41 for (size_t i = 0; i < num_threads; i++)
42 {
43 m_workers.emplace_back(std::thread(&ThreadPool::thread_fn, this, i));
44 }
45 }
46
47 size_t get_num_threads() const
48 {
49 return m_workers.size();
50 }
51
52 virtual ~ThreadPool()
53 {
54 // std::cerr << "-> ~ThreadPool" << std::endl;
55 stop();
56 // std::cerr << "<- ~ThreadPool" << std::endl;
57 }
58
59 void stop()
60 {
61 m_stop = true;
62 // push null jobs to ensure having enough joinable jobs
63 for (size_t i = 0 ; i < m_workers.size(); i++)
64 {
65 push(nullptr);
66 }
67
68 for (auto& t : m_workers)
69 {
70 if (t.joinable())
71 {
72 t.join();
73 }
74 }
75
76 while (!m_workers.empty())
77 {
78 if (!m_workers.back().joinable())
79 {
80 m_workers.pop_back();
81 }
82 else
83 {
84 throw std::runtime_error("All workers must be unjoinable (inactive threads) here.");
85 }
86 }
87 }
88
89 size_t running()
90 {
91 return m_workers.size();
92 }
93
94 inline void push(void* job)
95 {
96 std::lock_guard<std::mutex> l(m_mutex);
97 m_jobs.push(job);
98 m_cv.notify_one();
99 }
100
101protected:
102
106 inline bool wait_for_new_job(void** job)
107 {
108 std::unique_lock<std::mutex> l(m_mutex);
109 m_cv.wait(l, [this](){ return !m_jobs.empty() || m_stop; });
110
111 if (!m_jobs.empty())
112 {
113 *job = m_jobs.front();
114 m_jobs.pop();
115
116 return true;
117 }
118
119 return false;
120 }
121
122 inline void wait_for_any_job_notification(const std::function<bool()> fn)
123 {
124 std::unique_lock<std::mutex> l(m_mutex_notify);
125 m_cv_notify.wait(l, [&fn](){ return fn(); });
126 }
127
128 void thread_fn(size_t worker_id)
129 {
130 void* job = nullptr;
131 // loop to dispatch pushed jobs to the threads of this pool
132 // wait_for_new_job is blocking until a job becomes available
133 while (true)
134 {
135 // std::cerr << "thread_fn " << worker_id << " main loop" << std::endl;
136 if (wait_for_new_job(&job))
137 {
138 // std::cerr << "wait_for_new_job is true" << std::endl;
139 if (nullptr == job)
140 {
141 // we should get a null job only when stopping
142 // std::cerr << "wait_for_new_job: we should get a null job only when stopping" << std::endl;
143 break;
144 }
145 // std::cerr << "worker: running job " << (void*) job << std::endl;
146 P::run_one_job(static_cast<P*>(this), worker_id, job);
147 // std::cerr << "worker: completed job " << (void*) job << std::endl;
148 m_cv_notify.notify_all();
149 // std::cerr << "notify_all done" << std::endl;
150 }
151 else
152 {
153 break;
154 }
155 }
156
157 if (!m_stop)
158 {
159 throw std::runtime_error("Worker finished but stop flag isn't set.");
160 }
161 // std::cerr << "thread_fn done" << std::endl;
162 }
163
164 std::vector<std::thread> m_workers;
165 std::atomic<bool> m_stop;
166
167 std::queue<void*> m_jobs;
168 std::mutex m_mutex;
169 std::condition_variable m_cv;
170
171 std::mutex m_mutex_notify;
172 std::condition_variable m_cv_notify;
173};
174
175} // namespace deeplima
176
177#endif
bool wait_for_new_job(void **job)
This will wait until a job is available and then job parameter will be set to this available which wi...
void init(size_t num_threads)
Definition thread_pool.h:35
std::mutex m_mutex_notify
void push(void *job)
Definition thread_pool.h:94
void wait_for_any_job_notification(const std::function< bool()> fn)
size_t get_num_threads() const
Definition thread_pool.h:47
ThreadPool(size_t num_threads=0)
Definition thread_pool.h:26
std::queue< void * > m_jobs
std::atomic< bool > m_stop
std::vector< std::thread > m_workers
std::condition_variable m_cv_notify
void thread_fn(size_t worker_id)
std::condition_variable m_cv