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

如何并行化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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.07 11:42:50