matxscript运行时系统之线程池系统

鱿鱼圈 Lv4

仓库链接:bytedance/matxscript: A high-performance, extensible Python AOT compiler.

1 概述

线程池系统是matxscript运行时的重要组成部分,负责提供并发执行能力。该系统实现了基于锁和无锁两种线程池模式,通过统一的执行器(ThreadPoolExecutor)对外暴露接口,支持任务并行处理、异步调用等功能。

线程池系统位于运行时层的核心位置,为上层应用提供高性能的并发执行支持。它被广泛应用于各种需要并发处理的场景,如批量数据处理、并行计算等。

2 整体架构

从宏观角度来看:线程池执行器–依赖–>线程池–依赖–>任务队列

具体而言,线程池系统的整体架构采用分层设计模式,包含以下几个核心组件:

  1. ThreadPoolExecutor(线程池执行器):对外提供统一接口,封装不同类型的线程池实现
  2. IThreadPool(线程池接口):定义线程池的基本行为规范
  3. LockBasedThreadPool(基于锁的线程池):使用互斥锁实现的任务队列管理
  4. LockFreeThreadPool(无锁线程池):基于无锁队列实现的高性能线程池
  5. SPSCLockFreeThreadPool(单生产者单消费者无锁线程池):针对特定场景优化的无锁实现
  6. 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将一个函数应用到输入列表或元组的每个元素上,并返回结果列表或元组。

重载版本:

  1. ParallelFor(op, inputs) - 使用默认线程数和组大小
  2. ParallelFor(op, inputs, expt_num_threads, group_size) - 指定线程数和组大小

参数说明:

  • op: 要应用的函数对象
  • inputs: 输入数据列表或元组
  • expt_num_threads: 期望使用的线程数,默认为线程池大小+1
  • group_size: 任务分组大小,默认为1

ParallelStarMap 并行星形映射

ParallelStarMap类似于ParallelFor,但会将输入的每个元素作为参数列表展开传递给函数。

重载版本:

  1. ParallelStarMap(op, inputs) - 使用默认线程数和组大小
  2. 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
// Copyright 2022 ByteDance Ltd. and/or its affiliates.
/*
* Acknowledgement:
* Taken from http://www.1024cores.net/home/lock-free-algorithms/queues/bounded-mpmc-queue
*
* Licensed to the Apache Software Foundation (ASF) under one
* or more contributor license agreements. See the NOTICE file
* distributed with this work for additional information
* regarding copyright ownership. The ASF licenses this file
* to you under the Apache License, Version 2.0 (the
* "License"); you may not use this file except in compliance
* with the License. You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing,
* software distributed under the License is distributed on an
* "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
* KIND, either express or implied. See the License for the
* specific language governing permissions and limitations
* under the License.
*/
#pragma once

#include <cstddef>
#include <cstdint>

#include <atomic>
#include <memory>

namespace matxscript {
namespace runtime {

template <typename T>
class MPMCBoundedQueue {
public:
/**
* buffer_size must be 2^n
* @param buffer_size
*/
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;
}

// problem: heavy enqueue competition, need better implementation in the future
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;
};

} // namespace runtime
} // namespace matxscript

类属性:

  • 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) {
// 1. 读取当前入队位置
size_t pos = enqueue_pos_.load(std::memory_order_relaxed);
for (;;) {
// 2. 找到对应的槽位
cell_t* cell = &buffer_[pos & buffer_mask_];

// 3. 读取该槽位的序列号
size_t seq = cell->sequence_.load(std::memory_order_acquire);

// 4. 计算差异
intptr_t dif = (intptr_t)seq - (intptr_t)pos;

if (dif == 0) { // 槽位为空,可以写入
// 5. CAS 更新入队位置(防止其他线程同时写入)
if (enqueue_pos_.compare_exchange_weak(pos, pos + 1,
std::memory_order_relaxed))
break; // 成功获取槽位
} else if (dif < 0) {
// 6. 队列已满(消费者还未取走数据)
return false;
} else {
// 7. 其他线程正在操作,重新尝试
pos = enqueue_pos_.load(std::memory_order_relaxed);
}
}

// 8. 写入数据
cell->data_ = std::forward<U>(data);

// 9. 发布数据:更新序列号为 pos + 1
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) {
// 1. 读取当前出队位置
size_t pos = dequeue_pos_.load(std::memory_order_relaxed);
for (;;) {
// 2. 找到对应的槽位
cell_t* cell = &buffer_[pos & buffer_mask_];

// 3. 读取序列号
size_t seq = cell->sequence_.load(std::memory_order_acquire);

// 4. 计算差异(期待 seq = pos + 1)
intptr_t dif = (intptr_t)seq - (intptr_t)(pos + 1);

if (dif == 0) { // 有数据,可以读取
// 5. CAS 更新出队位置
if (dequeue_pos_.compare_exchange_weak(pos, pos + 1,
std::memory_order_relaxed))
break; // 成功获取槽位
} else if (dif < 0) {
// 6. 队列为空(生产者还未写入数据)
return false;
} else {
// 7. 其他线程正在操作,重新尝试
pos = dequeue_pos_.load(std::memory_order_relaxed);
}
}

// 8. 读取数据
data = std::move(cell->data_);

// 9. 标记槽位为空:更新序列号为 pos + buffer_mask_ + 1
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)

  • 是什么:缓冲区中的一个存储单元
  • 组成:每个槽位包含两部分:
  1. 序列号(sequence_):原子计数器,用于同步
  2. 数据(data_):实际存储的元素
  • 位置:由索引决定在缓冲区中的位置

🔄 序列号(Sequence Number)

  • 是什么:每个槽位的"状态标签"
  • 作用
  1. 标识槽位的可用状态
  2. 实现生产者和消费者的同步
  3. 避免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:
// need to keep track of threads so we can join them
std::vector<std::thread> workers_;
// the task queue
MPMCBoundedQueue<IRunnablePtr> tasks_;
// stop flag
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_; //工作的线程池
// the task queue
MPMCBoundedQueue<IRunnablePtr> tasks_; // 任务队列
// stop flag
bool stop_ = false;
std::string name_;
int64_t intervals_ns_;
pid_t belong_to_pid_; // 属于的进程id

工作线程创建机制

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
// the constructor just launches some amount of workers
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) {
// 设置线程名称(Linux平台)
#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; // 设置停止标志

// 只有在创建线程池的进程中才join线程
if (cur_pid == belong_to_pid_) {
for (std::thread& worker : workers_) {
if (worker.joinable()) {
worker.join();
}
}
} else {
// fork后的子进程中detach线程
for (std::thread& worker : workers_) {
worker.detach();
}
}
}

停止机制特点

  • 通过原子或线程安全的方式设置 stop_ 标志
  • 工作线程在循环中检查该标志,决定是否退出
  • 析构函数中join所有工作线程,确保资源正确释放
  • 处理fork场景,避免在子进程中操作无效的线程

任务执行流程

无论是无锁还是基于锁的实现,任务执行都遵循统一流程:

  1. 工作线程从队列获取 IRunnable 任务对象
  2. 调用 task->Run() 方法执行任务
  3. Run() 方法内部调用 RunImpl() 执行实际计算逻辑
  4. 捕获并存储执行过程中的异常
  5. 任务完成后调用 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; // 默认每组任务大小为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;
//设计两种步长(step_r/step_l)确保任务均匀分配到所有线程

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 {
// 创建不需要解包参数的LockFreeRunnable任务
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 {
// 创建基于锁的Runnable任务
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()) {
// fix nested pmap: 如果当前线程已经是线程池线程,顺序执行所有任务
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);
}

  • 标题: matxscript运行时系统之线程池系统
  • 作者: 鱿鱼圈
  • 创建于 : 2025-11-03 23:50:00
  • 更新于 : 2026-06-12 17:14:55
  • 链接: https://yuyanqi.com/2025/11/03/matxscript线程池/
  • 版权声明: 本文章采用 CC BY-NC-SA 4.0 进行许可。
评论