仓库链接:bytedance/matxscript: A high-performance, extensible Python AOT compiler.
1 概述
线程池系统是matxscript运行时的重要组成部分,负责提供并发执行能力。该系统实现了基于锁和无锁两种线程池模式,通过统一的执行器(ThreadPoolExecutor)对外暴露接口,支持任务并行处理、异步调用等功能。
线程池系统位于运行时层的核心位置,为上层应用提供高性能的并发执行支持。它被广泛应用于各种需要并发处理的场景,如批量数据处理、并行计算等。
2 整体架构
从宏观角度来看:线程池执行器–依赖–>线程池–依赖–>任务队列
具体而言,线程池系统的整体架构采用分层设计模式,包含以下几个核心组件:
ThreadPoolExecutor(线程池执行器) :对外提供统一接口,封装不同类型的线程池实现
IThreadPool(线程池接口) :定义线程池的基本行为规范
LockBasedThreadPool(基于锁的线程池) :使用互斥锁实现的任务队列管理
LockFreeThreadPool(无锁线程池) :基于无锁队列实现的高性能线程池
SPSCLockFreeThreadPool(单生产者单消费者无锁线程池) :针对特定场景优化的无锁实现
IRunnable(可运行任务接口) :定义任务的基本行为规范
该架构图展示了线程池系统的主要组件及其关系。ThreadPoolExecutor作为统一入口,通过IThreadPool接口与具体实现解耦,支持多种线程池类型的选择。IRunnable接口定义了任务的基本行为,不同线程池实现对应不同的任务类型。
3 核心类解析
ThreadPoolExecutor 线程池执行器
ThreadPoolExecutor是线程池系统的统一入口,负责管理线程池实例并提供对外接口。其主要职责包括:
封装底层线程池实现
提供ParallelFor、ParallelStarMap等并行处理接口
支持异步任务提交(Submit/ApplyAsync)
处理嵌套调用情况下的特殊逻辑
关键成员变量:
lock_free_: 标识是否使用无锁线程池
thread_num_: 线程数量
pool_: 底层线程池实例
serial_: 用于生成任务序列号的原子计数器
pool_thread_ids_: 记录线程池中所有线程ID的集合
IThreadPool 线程池接口
IThreadPool定义了线程池的基本行为规范,所有具体实现都需要继承此接口:
Enqueue: 添加单个任务到线程池
EnqueueBulk: 批量添加任务到线程池
GetThreadsNum: 获取线程池中的线程数量
GetThreadIds: 获取线程池中所有线程的ID
WaitBulk: 等待一批任务完成
LockBasedThreadPool 基于锁的线程池
LockBasedThreadPool使用传统的互斥锁和条件变量实现任务队列管理:
使用std::mutex保护任务队列
使用std::condition_variable进行线程间通信
任务队列为std::queue<IRunnablePtr>类型
线程函数循环等待并执行任务
LockFreeThreadPool 无锁线程池
LockFreeThreadPool基于无锁队列实现,具有更高的并发性能:
使用MPMCBoundedQueue作为任务队列
通过原子操作实现无锁入队和出队
线程忙等待而非阻塞休眠
支持设置轮询间隔参数
SPSCLockFreeThreadPool 单生产者单消费者无锁线程池
SPSCLockFreeThreadPool是对LockFreeThreadPool的进一步封装:
内部维护多个单线程的LockFreeThreadPool实例
根据任务序列号分配到不同线程池实现负载均衡
提供更好的缓存局部性和更低的竞争开销
4 接口详解
ThreadPoolExecutor提供了丰富的接口用于并发任务处理:
ParallelFor 并行映射
ParallelFor将一个函数应用到输入列表或元组的每个元素上,并返回结果列表或元组。
重载版本:
ParallelFor(op, inputs) - 使用默认线程数和组大小
ParallelFor(op, inputs, expt_num_threads, group_size) - 指定线程数和组大小
参数说明:
op: 要应用的函数对象
inputs: 输入数据列表或元组
expt_num_threads: 期望使用的线程数,默认为线程池大小+1
group_size: 任务分组大小,默认为1
ParallelStarMap 并行星形映射
ParallelStarMap类似于ParallelFor,但会将输入的每个元素作为参数列表展开传递给函数。
重载版本:
ParallelStarMap(op, inputs) - 使用默认线程数和组大小
ParallelStarMap(op, inputs, expt_num_threads, group_size) - 指定线程数和组大小
Submit 异步任务提交
Submit方法允许异步提交任务并返回Future对象用于获取结果。
函数签名:RTValue Submit(PyArgs args)
参数说明:
args[0]: 要执行的可调用对象
args[1..n]: 传递给可调用对象的参数
ApplyAsync 异步调用
ApplyAsync是Submit方法的底层实现,直接接受函数对象和参数。
函数签名:RTValue ApplyAsync(const UserDataRef& op, const PyArgs& args)
5 代码实现分析
任务队列
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 #pragma once #include <cstddef> #include <cstdint> #include <atomic> #include <memory> namespace matxscript { namespace runtime { template <typename T> class MPMCBoundedQueue { public : MPMCBoundedQueue (size_t buffer_size) : buffer_ (new cell_t [buffer_size]), buffer_mask_ (buffer_size - 1 ) { if (!((buffer_size >= 2 ) && ((buffer_size & (buffer_size - 1 )) == 0 ))) { abort (); } for (size_t i = 0 ; i != buffer_size; i += 1 ) buffer_[i].sequence_.store (i, std::memory_order_relaxed); enqueue_pos_.store (0 , std::memory_order_relaxed); dequeue_pos_.store (0 , std::memory_order_relaxed); } virtual ~MPMCBoundedQueue () { delete [] buffer_; } template <class U > bool enqueue (U&& data) { cell_t * cell; size_t pos = enqueue_pos_.load (std::memory_order_relaxed); for (;;) { cell = &buffer_[pos & buffer_mask_]; size_t seq = cell->sequence_.load (std::memory_order_acquire); intptr_t dif = (intptr_t )seq - (intptr_t )pos; if (dif == 0 ) { if (enqueue_pos_.compare_exchange_weak (pos, pos + 1 , std::memory_order_relaxed)) break ; } else if (dif < 0 ) return false ; else pos = enqueue_pos_.load (std::memory_order_relaxed); } cell->data_ = std::forward<U>(data); cell->sequence_.store (pos + 1 , std::memory_order_release); return true ; } template <class U > bool try_enqueue (U&& data) { cell_t * cell; size_t pos = enqueue_pos_.load (std::memory_order_relaxed); cell = &buffer_[pos & buffer_mask_]; size_t seq = cell->sequence_.load (std::memory_order_acquire); intptr_t dif = (intptr_t )seq - (intptr_t )pos; if (dif == 0 ) { if (enqueue_pos_.compare_exchange_weak (pos, pos + 1 , std::memory_order_relaxed)) { cell->data_ = std::forward<U>(data); cell->sequence_.store (pos + 1 , std::memory_order_release); return true ; } } return false ; } template <class U > bool enqueue_bulk (U* data, size_t size) { cell_t * cell; size_t pos = enqueue_pos_.load (std::memory_order_relaxed); for (;;) { cell = &buffer_[pos & buffer_mask_]; size_t seq = cell->sequence_.load (std::memory_order_acquire); intptr_t dif = (intptr_t )seq - (intptr_t )pos; if (dif == 0 ) { if (enqueue_pos_.compare_exchange_weak (pos, pos + size, std::memory_order_relaxed)) break ; } else if (dif < 0 ) return false ; else pos = enqueue_pos_.load (std::memory_order_relaxed); } for (size_t i = 0 ; i < size; ++i) { cell_t * cell = &buffer_[(pos + i) & buffer_mask_]; cell->data_ = data[i]; cell->sequence_.store (pos + 1 + i, std::memory_order_release); } return true ; } bool dequeue (T& data) { cell_t * cell; size_t pos = dequeue_pos_.load (std::memory_order_relaxed); for (;;) { cell = &buffer_[pos & buffer_mask_]; size_t seq = cell->sequence_.load (std::memory_order_acquire); intptr_t dif = (intptr_t )seq - (intptr_t )(pos + 1 ); if (dif == 0 ) { if (dequeue_pos_.compare_exchange_weak (pos, pos + 1 , std::memory_order_relaxed)) break ; } else if (dif < 0 ) return false ; else pos = dequeue_pos_.load (std::memory_order_relaxed); } data = std::move (cell->data_); cell->sequence_.store (pos + buffer_mask_ + 1 , std::memory_order_release); return true ; } bool try_dequeue (T& data) { cell_t * cell; size_t pos = dequeue_pos_.load (std::memory_order_relaxed); cell = &buffer_[pos & buffer_mask_]; size_t seq = cell->sequence_.load (std::memory_order_acquire); intptr_t dif = (intptr_t )seq - (intptr_t )(pos + 1 ); if (dif == 0 ) { if (dequeue_pos_.compare_exchange_weak (pos, pos + 1 , std::memory_order_relaxed)) { data = std::move (cell->data_); cell->sequence_.store (pos + buffer_mask_ + 1 , std::memory_order_release); return true ; } } return false ; } inline size_t size () const { return enqueue_pos_ - dequeue_pos_; } inline bool empty () { return size () == 0 ; } protected : struct cell_t { std::atomic<size_t > sequence_; T data_; }; static size_t const cacheline_size = 64 ; typedef char cacheline_pad_t [cacheline_size]; cacheline_pad_t pad0_; cell_t * const buffer_; size_t const buffer_mask_; cacheline_pad_t pad1_; std::atomic<size_t > enqueue_pos_; cacheline_pad_t pad2_; std::atomic<size_t > dequeue_pos_; cacheline_pad_t pad3_; MPMCBoundedQueue (MPMCBoundedQueue const &) = delete ; void operator =(MPMCBoundedQueue const &) = delete ; }; } }
类属性:
cell_t: 槽位,每个槽位包含:
sequence_:序列号,用于同步
data_:存储的数据
cell_t* const buffer_;: 环形缓冲区,缓冲区大小必须是 2 的幂次方(2^n)
size_t const buffer_mask_;: 掩码(用于取模运算)使用掩码 & buffer_mask_ 代替取模 % buffer_size,效率更高
std::atomic<size_t> enqueue_pos_;: 入队位置
std::atomic<size_t> dequeue_pos_;: 出队位置
cacheline_pad_t * : 缓存行对齐,防止伪共享(false sharing)
入队算法
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 bool enqueue (U&& data) { size_t pos = enqueue_pos_.load (std::memory_order_relaxed); for (;;) { cell_t * cell = &buffer_[pos & buffer_mask_]; size_t seq = cell->sequence_.load (std::memory_order_acquire); intptr_t dif = (intptr_t )seq - (intptr_t )pos; if (dif == 0 ) { if (enqueue_pos_.compare_exchange_weak (pos, pos + 1 , std::memory_order_relaxed)) break ; } else if (dif < 0 ) { return false ; } else { pos = enqueue_pos_.load (std::memory_order_relaxed); } } cell->data_ = std::forward<U>(data); cell->sequence_.store (pos + 1 , std::memory_order_release); return true ; }
出队算法
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 bool dequeue (T& data) { size_t pos = dequeue_pos_.load (std::memory_order_relaxed); for (;;) { cell_t * cell = &buffer_[pos & buffer_mask_]; size_t seq = cell->sequence_.load (std::memory_order_acquire); intptr_t dif = (intptr_t )seq - (intptr_t )(pos + 1 ); if (dif == 0 ) { if (dequeue_pos_.compare_exchange_weak (pos, pos + 1 , std::memory_order_relaxed)) break ; } else if (dif < 0 ) { return false ; } else { pos = dequeue_pos_.load (std::memory_order_relaxed); } } data = std::move (cell->data_); cell->sequence_.store (pos + buffer_mask_ + 1 , std::memory_order_release); return true ; }
疑惑解答
📦 缓冲区(Buffer)
是什么 :一块连续的内存区域,用于存储数据元素
大小 :必须是 2 的幂次方(如 2, 4, 8, 16, 32…)
作用 :队列的物理存储空间
🔢 索引(Index)
是什么 :访问缓冲区中具体位置的数字(0, 1, 2, 3…)
计算 :索引 = 位置 & 缓冲区掩码
示例 :位置=5,掩码=3,索引=1(因为 5 & 3 = 1)
🎯 槽位(Slot/Cell)
是什么 :缓冲区中的一个存储单元
组成 :每个槽位包含两部分:
序列号(sequence_) :原子计数器,用于同步
数据(data_) :实际存储的元素
🔄 序列号(Sequence Number)
标识槽位的可用状态
实现生产者和消费者的同步
避免ABA问题
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 ┌───────────────────────────────────────────────────┐ │ 缓冲区(Buffer) │ │ 大小为4 ,连续内存块,存储4 个槽位 │ ├────────────┬────────────┬────────────┬────────────┤ │ 槽位[0 ] │ 槽位[1 ] │ 槽位[2 ] │ 槽位[3 ] │ │ 索引=0 │ 索引=1 │ 索引=2 │ 索引=3 │ ├────────────┼────────────┼────────────┼────────────┤ │ sequence=8 │ sequence=6 │ sequence=7 │ sequence=8 │ │ data=null │ data=B │ data=C │ data=D │ └────────────┴────────────┴────────────┴────────────┘ ↑ ↑ ↑ ↑ │ │ │ │ 状态:空 状态:满 状态:满 状态:满 (可写入) (待消费) (待消费) (待消费) enqueue_pos = 9 (下一个要写入的"逻辑位置" ) dequeue_pos = 6 (下一个要读取的"逻辑位置" ) buffer_mask = 3 (缓冲区大小-1 )
如何从逻辑位置找到物理槽位?
关键映射公式:
1 2 3 4 逻辑位置(pos) → 物理索引(index) → 槽位(cell) index = pos & buffer_mask; cell = &buffer_[index];
示例计算:
假设 buffer_size = 4, buffer_mask = 3
1 2 3 逻辑位置 pos = 9 索引 index = 9 & 3 = 1 (二进制: 1001 & 0011 = 0001) 槽位 cell = buffer_[1]
无锁线程池
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 class LockFreeRunnable : public IRunnable { public : bool Done () override { return finish_; } protected : void SetDone () override { finish_ = true ; } private : volatile bool finish_ = false ; friend class LockFreeThreadPool ; }; class LockFreeThreadPool : public IThreadPool { public : explicit LockFreeThreadPool (size_t threads, const std::string& name, int64_t intervals_ns) ; explicit LockFreeThreadPool (size_t threads, const std::string& name) ; ~LockFreeThreadPool () override ; void Enqueue (IRunnablePtr& runner, size_t seq) override ; void EnqueueBulk (std::vector<IRunnablePtr>& runners) override ; size_t GetThreadsNum () const override ; std::vector<std::thread::id> GetThreadIds () const override ; protected : static void ThreadEntry (LockFreeThreadPool* pool, const std::string& name) ; private : std::vector<std::thread> workers_; MPMCBoundedQueue<IRunnablePtr> tasks_; bool stop_ = false ; std::string name_; int64_t intervals_ns_; pid_t belong_to_pid_; };
类属性
1 2 3 4 5 6 7 8 std::vector<std::thread> workers_; MPMCBoundedQueue<IRunnablePtr> tasks_; bool stop_ = false ;std::string name_; int64_t intervals_ns_;pid_t belong_to_pid_;
工作线程创建机制
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 LockFreeThreadPool::LockFreeThreadPool (size_t threads, const std::string& name, int64_t intervals_ns) : stop_ (false ), name_ (name), tasks_ (4096 ), intervals_ns_ (intervals_ns) { #ifdef _WIN32 belong_to_pid_ = GetCurrentProcessId (); #else belong_to_pid_ = getpid (); #endif for (size_t i = 0 ; i < threads; ++i) { char buffer[16 ] = {0 }; snprintf (buffer, sizeof (buffer), "T%zu.%s" , i, name.c_str ()); workers_.emplace_back (LockFreeThreadPool::ThreadEntry, this , std::string (buffer)); } }
关键特点 :
使用 std::vector<std::thread> 存储工作线程
每个线程都执行静态成员函数 ThreadEntry 作为入口点
为线程分配唯一名称,便于调试和监控
记录所属进程ID,确保在fork场景下安全处理
1 workers_.emplace_back (ThreadEntry, this , name + "_T" + std::to_string (i));
这行代码会:
在 workers_(std::vector<std::thread>)中创建并初始化 一个新的线程对象
以 ThreadEntry 函数作为线程的入口点
传递 this(线程池对象指针)和线程名称作为参数
线程立即开始执行 ThreadEntry 函数
任务队列监听机制
工作线程通过 ThreadEntry 函数持续监听任务队列
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 void LockFreeThreadPool::ThreadEntry (LockFreeThreadPool* pool, const std::string& name) { #ifdef __linux__ pthread_setname_np (pthread_self (), name.c_str ()); #endif int64_t sleep_intervals_ns = pool->intervals_ns_; if (sleep_intervals_ns <= 0 ) { sleep_intervals_ns = 1 ; } for (;;) { IRunnablePtr task = nullptr ; for (;;) { if (pool->tasks_.try_dequeue (task)) { break ; } if (pool->stop_) { break ; } std::this_thread::sleep_for (std::chrono::nanoseconds (1 )); } if (pool->stop_) { return ; } else if (task != nullptr ) { task->Run (); } } }
无锁监听特点 :
使用 try_dequeue 非阻塞方式获取任务
空队列时采用自旋等待(busy-waiting)策略
通过 sleep_for(1ns) 降低CPU占用
响应速度快,但可能消耗更多CPU资源
线程池停止与资源释放
线程池都通过 stop_ 标志控制线程生命周期:
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 LockFreeThreadPool::~LockFreeThreadPool () { #ifdef _WIN32 auto cur_pid = GetCurrentProcessId (); #else auto cur_pid = getpid (); #endif stop_ = true ; if (cur_pid == belong_to_pid_) { for (std::thread& worker : workers_) { if (worker.joinable ()) { worker.join (); } } } else { for (std::thread& worker : workers_) { worker.detach (); } } }
停止机制特点 :
通过原子或线程安全的方式设置 stop_ 标志
工作线程在循环中检查该标志,决定是否退出
析构函数中join所有工作线程,确保资源正确释放
处理fork场景,避免在子进程中操作无效的线程
任务执行流程
无论是无锁还是基于锁的实现,任务执行都遵循统一流程:
工作线程从队列获取 IRunnable 任务对象
调用 task->Run() 方法执行任务
Run() 方法内部调用 RunImpl() 执行实际计算逻辑
捕获并存储执行过程中的异常
任务完成后调用 SetDone() 更新状态
线程池执行器
职责:
封装底层线程池实现
提供ParallelFor、ParallelStarMap等并行处理接口
创建任务并提交到任务队列
支持异步任务提交(Submit/ApplyAsync)
处理嵌套调用情况下的特殊逻辑
关键成员变量:
lock_free_: 标识是否使用无锁线程池
thread_num_: 线程数量
pool_: 底层线程池实例
serial_: 用于生成任务序列号的原子计数器
pool_thread_ids_: 记录线程池中所有线程ID的集合
并行任务的具体定义
任务的具体执行逻辑
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 template <typename RunnableType, bool UnpackArgs = false >class ParallelForTask : public RunnableType { public : ParallelForTask (const UserDataRef& op, const Any* input_first, RTValue* output_first, int64_t len) : op_ (&op), input_first_ (input_first), input_last_ (input_first + len), output_first_ (output_first) { } void RunImpl () override { while (input_first_ != input_last_) { if (UnpackArgs) { switch (input_first_->type_code ()) { case TypeIndex::kRuntimeList: { auto args = input_first_->template AsObjectRefNoCheck <List>(); *output_first_ = op_->generic_call (PyArgs (args.data (), args.size ())); } break ; case TypeIndex::kRuntimeTuple: { auto args = input_first_->template AsObjectRefNoCheck <Tuple>(); *output_first_ = op_->generic_call (PyArgs (args.begin (), args.size ())); } break ; case TypeIndex::kRuntimeFTList: { auto num_args = kernel_object___len__ (*input_first_); Iterator iterable = Kernel_Iterable::make (*input_first_); std::vector<RTValue> args; args.reserve (num_args); bool has_next = iterable.HasNext (); while (has_next) { args.emplace_back (iterable.Next (&has_next)); } *output_first_ = op_->generic_call (PyArgs (args.data (), args.size ())); } break ; default : { MXTHROW << "matx.pstarmap(f, iterable) expect iterable[i] is list or tuple, but get " << input_first_->type_name (); } break ; } } else { *output_first_ = op_->generic_call (PyArgs (input_first_, 1 )); } ++input_first_; ++output_first_; } } private : const UserDataRef* op_; const Any* input_first_; const Any* input_last_; RTValue* output_first_; };
并行任务的创建、管理、调度
这个函数是整个线程池执行器的核心,体现了高性能并行计算的设计理念,通过合理的任务分配、线程管理和资源利用,实现了高效的并行数据处理能力。
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 void ThreadPoolExecutor::ParallelForImpl (const UserDataRef& op, const Any* inputs_begin,const Any* inputs_end,int64_t expt_num_threads,int64_t group_size,RTValue* outputs_begin, bool unpack_args) { int64_t input_size = inputs_end - inputs_begin; if (expt_num_threads <= 0 ) { expt_num_threads = thread_num_ + 1 ; } if (group_size <= 0 ) { group_size = 1 ; } MXCHECK (input_size % group_size == 0 ) << "Expect the number of tasks to be a multiple of " << group_size << ", but get " << input_size << "" ; int64_t num_group = input_size / group_size; int64_t step_r = group_size * ((num_group + expt_num_threads - 1 ) / expt_num_threads); int64_t step_l = step_r - group_size; int64_t pos = 0 ; int64_t step = step_r; bool need_change = true ; std::vector<internal::IRunnablePtr> tasks; tasks.reserve (expt_num_threads); for (int64_t i = 0 ; i < expt_num_threads && pos < input_size; ++i) { if (need_change && step_l != 0 && pos + step_l * (expt_num_threads - i) == input_size) { step = step_l; need_change = false ; } if (lock_free_) { if (unpack_args) { auto task = std::make_shared<ParallelForTask<internal::LockFreeRunnable, true >>( op, inputs_begin + pos, outputs_begin + pos, step); tasks.push_back (std::static_pointer_cast <internal::IRunnable>(task)); } else { auto task = std::make_shared<ParallelForTask<internal::LockFreeRunnable, false >>( op, inputs_begin + pos, outputs_begin + pos, step); tasks.push_back (std::static_pointer_cast <internal::IRunnable>(task)); } } else { if (unpack_args) { auto task = std::make_shared<ParallelForTask<internal::LockBasedRunnable, true >>( op, inputs_begin + pos, outputs_begin + pos, step); tasks.push_back (std::static_pointer_cast <internal::IRunnable>(task)); } else { auto task = std::make_shared<ParallelForTask<internal::LockBasedRunnable, false >>( op, inputs_begin + pos, outputs_begin + pos, step); tasks.push_back (std::static_pointer_cast <internal::IRunnable>(task)); } } pos += step; } auto cur_tid = std::this_thread::get_id (); if (pool_thread_ids_.find (cur_tid) != pool_thread_ids_.end ()) { for (auto & task : tasks) { task->Run (); } } else { size_t task_size = tasks.size (); if (task_size > 1 ) { size_t seq = serial_.fetch_add (task_size - 1 , std::memory_order_relaxed); for (size_t i = 1 ; i < tasks.size (); ++i) { pool_->Enqueue (tasks[i], seq + i - 1 ); } } internal::IRunnablePtr& first_task = tasks[0 ]; first_task->Run (); } internal::IThreadPool::WaitBulk (tasks); }