PyTorch C++扩展:多线程下高效累加Tensor行结果的最优方案
PyTorch C++扩展并行优化问题与解决方案
问题背景
我正在实现一个带并行机制的PyTorch C++扩展,目标输出为N×D矩阵,设计为每个线程负责计算一行。目前采用OpenMP并行循环实现,但性能表现不佳,性能分析显示线程大量时间处于等待状态,推测是共享的result张量导致的同步开销。另外我认为共享BallTree指针不会影响性能,希望得到验证,同时寻求高效累加行结果的优化方案。
简化代码片段
at::Tensor nearestKKeys(at::Tensor queries, at::Tensor keys, int k, int maxLeafSize) { /* queries: (Nq, D) keys: (Nk, D) k: int return: (Nq, k) */ int Nq = queries.size(0); at::Tensor result = at::empty({Nq, k}, at::ScalarType::Int); at::Tensor indices = at::arange(Nk, keys.device()); BallTreePtr pBallTree = buildBallTree( keys, indices, maxLeafSize, 0 ); #pragma omp parallel for for (int n=0; n<Nq; n++) { at::Tensor query = queries.index({n}); BestMatchesPtr pBestMatches = std::make_shared<BestMatches>(k); // 将当前query的结果存入*pBestMatches pBallTree->query(query, query_norm, pBestMatches, k); // 将结果写入result的对应行 at::Tensor matches = pBestMatches->getMatches(); result.index_put_({n}, matches); } return result; }
核心疑问验证:共享BallTree指针的影响
BallTree在构建完成后仅被线程只读访问(query方法没有修改树结构的操作),多个线程同时读取只读数据不会产生竞争条件,也不需要任何同步机制。因此共享BallTreePtr确实不会导致线程等待或性能损耗,这部分的设计是合理的。
高效累加行结果的优化方案
1. 绕过高层张量API,直接操作底层内存
index_put_作为PyTorch的高层API,在并行写入时可能触发隐式的内存同步检查,导致线程阻塞。可以直接获取result的底层指针,通过内存拷贝完成写入:
int* result_ptr = result.data_ptr<int>(); // 在循环内: const int* matches_ptr = bestMatches.getMatches().data_ptr<int>(); std::copy(matches_ptr, matches_ptr + k, result_ptr + n * k);
2. 避免频繁创建小张量与智能指针
queries.index({n})会创建新的小张量,带来不必要的内存分配与拷贝开销。改为直接通过底层指针访问query数据:float* queries_ptr = queries.data_ptr<float>(); int D = queries.size(1); // 循环内: float* query_ptr = queries_ptr + n * D; at::Tensor query = at::from_blob(query_ptr, {D}, queries.options());- 替换
std::make_shared<BestMatches>为栈上分配对象,避免智能指针的原子操作开销:BestMatches bestMatches(k); pBallTree->query(query, query_norm, &bestMatches, k);
3. 优化OpenMP调度策略
默认的调度策略可能导致负载不均,根据计算复杂度调整调度方式:
- 若每个query的计算耗时相近,使用
schedule(static, chunk_size)(比如chunk_size设为64/128),减少线程调度开销; - 若计算耗时差异大,使用
schedule(dynamic)动态分配任务,避免部分线程闲置。
示例:
#pragma omp parallel for schedule(static, 64)
4. 线程本地存储复用
如果BestMatches的初始化开销较高,可以用OpenMP的threadprivate声明线程本地对象,循环内复用,避免重复初始化:
#pragma omp threadprivate(bestMatches) // 在并行区域外初始化每个线程的对象
优化后代码示例
at::Tensor nearestKKeys(at::Tensor queries, at::Tensor keys, int k, int maxLeafSize) { /* queries: (Nq, D) keys: (Nk, D) k: int return: (Nq, k) */ int Nq = queries.size(0); int D = queries.size(1); at::Tensor result = at::empty({Nq, k}, at::ScalarType::Int); int* result_ptr = result.data_ptr<int>(); float* queries_ptr = queries.data_ptr<float>(); at::Tensor indices = at::arange(Nk, keys.device()); BallTreePtr pBallTree = buildBallTree(keys, indices, maxLeafSize, 0); #pragma omp parallel for schedule(static, 64) for (int n=0; n<Nq; n++) { // 直接访问query底层数据 float* query_ptr = queries_ptr + n * D; at::Tensor query = at::from_blob(query_ptr, {D}, queries.options()); // 栈上分配BestMatches BestMatches bestMatches(k); pBallTree->query(query, query_norm, &bestMatches, k); // 直接写入result内存 const int* matches_ptr = bestMatches.getMatches().data_ptr<int>(); std::copy(matches_ptr, matches_ptr + k, result_ptr + n * k); } return result; }
内容的提问来源于stack exchange,提问作者Jonas De Schouwer
相关产品推荐
相关产品推荐

