You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.20 15:54:51