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
相关产品推荐
相关产品推荐

