如何并行化PyTorch中torch.linalg.solve()循环以提升运行速度?
加速torch.linalg.solve调用的优化方案
需求与现状
- 目标:尽可能提升调用
torch.linalg.solve()的函数运行速度 - 环境:拥有
input_array(尺寸50×100×100)和host_array(尺寸100×100×100),当前通过嵌套循环遍历input_array的行列,将solver函数处理input_array[:,i,j]的结果存入host_array[:,i,j] - 问题:当前运行速度缓慢,实际场景中单次函数调用耗时1秒,希望通过优化提速,确认并行化是否有效
原示例代码
import torch from tqdm import tqdm # 创建随机数据的host_array和input_array host_array = torch.zeros(100, 500, 500) input_array = torch.randn(50, 500, 500) # 创建虚拟系数矩阵A (50x100) A = torch.randn(50, 100) # 定义求解函数:处理input_array[:, i, j]并更新host_array[:, i, j] def solver(input_vector, A): # 求解线性方程组 solution = torch.linalg.solve(A.T@A, A.T@input_vector) return solution # 计算总运行次数 total_iterations = int(host_array.shape[1]*host_array.shape[2]) progress_bar = tqdm(total=total_iterations, dynamic_ncols=False, mininterval=1.0) # 遍历input_array处理 for i in range(host_array.shape[1]): for j in range(host_array.shape[2]): host_array[:,i,j] = solver(input_array[:,i,j], A) progress_bar.update(1)
优化方案
1. 向量化批量运算(核心优化,比手动并行高效)
原代码的嵌套循环存在大量Python解释器开销,PyTorch的核心优势是向量化批处理,可直接对整个数组做运算,避免循环:
- 首先注意到
solver中的计算等价于求最小二乘解,可利用矩阵广播特性批量处理所有输入 - 预计算
A.T@A和A.T,避免重复计算常量
优化后代码
import torch host_array = torch.zeros(100, 500, 500) input_array = torch.randn(50, 500, 500) A = torch.randn(50, 100) # 预计算常量,避免重复计算 ATA = A.T @ A AT = A.T # 批量处理所有输入:利用einsum完成批量矩阵乘法 AT_input = torch.einsum('ij,jhw->ihw', AT, input_array) # torch.linalg.solve支持批量输入,自动并行处理每个维度的求解 solution = torch.linalg.solve(ATA, AT_input) # 直接赋值给host_array host_array[:] = solution
2. GPU加速(若有可用GPU)
将所有张量移至GPU,PyTorch的CUDA实现会自动利用GPU多核并行计算,速度提升显著:
import torch # 检测并使用GPU device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') host_array = torch.zeros(100, 500, 500).to(device) input_array = torch.randn(50, 500, 500).to(device) A = torch.randn(50, 100).to(device) # 预计算常量 ATA = A.T @ A AT = A.T # 批量求解 AT_input = torch.einsum('ij,jhw->ihw', AT, input_array) solution = torch.linalg.solve(ATA, AT_input) host_array[:] = solution # 若需转回CPU host_array = host_array.cpu()
3. 关于手动并行化的说明
Python层面的多进程/多线程并行(如multiprocessing)反而可能因为GIL限制、张量数据拷贝开销,效果远不如PyTorch内置的向量化和GPU并行。torch.linalg.solve底层已做优化(CPU用OpenBLAS/MKL多线程,GPU用CUDA核并行),无需手动实现并行。
效果对比
原嵌套循环的速度瓶颈在于Python解释器开销,批量处理将计算移至PyTorch底层C++执行,速度可提升几十到上百倍;搭配GPU加速后,速度还能再提升一个数量级。
内容的提问来源于stack exchange,提问作者vdc
相关产品推荐
相关产品推荐

