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

PyTorch自定义矩阵运算代码并行化优化及Bus Error问题求助

问题描述

在深度学习工作中实现了如下PyTorch自定义矩阵运算代码,用于处理points和primitives两个张量:

def operations(points, primitives):
    """
    points shape: (batch size, number_of_points, 3)
    primitives shape: (batch_size,number_of_primitives,7)
    """
    gradient = torch.zeros(batch_size,number_of_points,number_of_primitives)
    for i in range(batch_size):
        
        temp_points = points[i,:,:]
        
        temp_primitives= primitives[i,:,:]
        temp = torch.zeros(number_of_points,number_of_primitives)
        for k in range(number_of_points):
            for j in range(number_of_primitives):
                temp[k,j] = torch.norm(temp_points[k,:]*temp_primitives[j,:3]+temp_primitives[j,3:6])
        gradient[i,:,:] = temp
    return gradient

该代码采用串行循环实现,效率较低,尝试了如下多线程实现:

def sstp(points,primitives):
    batch_size,number_of_points,_ = points.shape
    _,_,number_of_primitives = primitives.shape
    gradient = torch.zeros(batch_size,number_of_points,number_of_primitives)
    
    def level3(i,k,j):
        print("level 3 {} {}".format(k,j))
        temp_points = points[i,:,:]
        temp_primitives = primitives[i,:,:].transpose(1,0)
        gradient_ijk = torch.norm(temp_points[k,:]*temp_primitives[j,:3]+temp_primitives[j,3:6])
        gradient[i,k,j] = torch.norm(temp_points[k,:]*temp_primitives[j,:3]+temp_primitives[j,3:6])
    def level2(i):
        global pool
        pool.map(level3,[(i,k,j) for k in range(number_of_points) for j in range(number_of_primitives)])
    #level 1
    global pool
    pool = multiprocessing.pool.ThreadPool(100)
    pool.map(level2, range(batch_size))
    pool.close()
    return gradient

但运行时出现Bus Error错误,需要可行的并行化方案并解决该错误。


可行的并行化方案与错误修复

多线程实现报错原因

  • Bus Error源于内存访问冲突:PyTorch张量并非线程安全,多线程同时写入gradient张量会引发竞态条件,破坏内存结构。
  • Python的ThreadPool受GIL限制,无法真正利用多CPU核心,反而因线程切换和内存竞争拖慢性能,完全没必要用这种方式处理PyTorch运算。
  • PyTorch原生支持张量向量化运算,能自动利用GPU/CPU的并行计算能力,效率远高于手动循环或多线程。

向量化并行实现(推荐)

利用PyTorch的广播机制,将三重循环转化为批量张量运算,代码如下:

def operations_vectorized(points, primitives):
    """
    points shape: (batch_size, num_points, 3)
    primitives shape: (batch_size, num_primitives, 7)
    return shape: (batch_size, num_points, num_primitives)
    """
    batch_size, num_points, _ = points.shape
    _, num_primitives, _ = primitives.shape
    
    # 拆分primitives的缩放项和偏移项
    prim_scale = primitives[:, :, :3]  # (batch_size, num_primitives, 3)
    prim_offset = primitives[:, :, 3:6]  # (batch_size, num_primitives, 3)
    
    # 通过广播实现批量元素运算:自动对齐维度完成逐元素计算
    scaled_points = points.unsqueeze(2) * prim_scale.unsqueeze(1)  # (batch, num_points, num_primitives, 3)
    added = scaled_points + prim_offset.unsqueeze(1)  # (batch, num_points, num_primitives, 3)
    
    # 计算最后一维的L2范数,得到目标结果
    gradient = torch.norm(added, dim=-1)  # (batch, num_points, num_primitives)
    
    return gradient

效果说明

  • 无显式循环,完全利用PyTorch的并行计算能力,GPU上能发挥最大性能,CPU上也会自动用SIMD指令加速。
  • 规避了多线程的内存竞争问题,不会出现Bus Error。
  • 代码简洁,逻辑与原串行实现完全一致,计算结果无差异。

内容的提问来源于stack exchange,提问作者Bill

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.06 19:46:01