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

如何在PyTorch中并行计算多个N×N矩阵的连乘积?

矩阵连乘积的并行分组计算实现(PyTorch)

核心思路

用分治策略:把矩阵列表两两分组,并行计算每组内的矩阵乘积,得到新的矩阵列表后重复该过程,直到列表只剩一个矩阵,就是最终的连乘积。这种方式能充分利用PyTorch的GPU并行计算能力,比逐次迭代相乘效率更高。

具体实现方法

方法一:递归分治实现(逻辑清晰)

适合任意数量的矩阵(包括奇数个的情况),递归逻辑容易理解:

import torch

def parallel_mat_product(matrices):
    # 基线:只剩一个矩阵直接返回
    if len(matrices) == 1:
        return matrices[0]
    
    grouped = []
    idx = 0
    while idx < len(matrices):
        # 两两配对计算乘积
        if idx + 1 < len(matrices):
            grouped.append(torch.matmul(matrices[idx], matrices[idx+1]))
        else:
            # 奇数个矩阵时,最后一个单独保留到下一轮
            grouped.append(matrices[idx])
        idx += 2
    
    # 递归处理分组后的新列表
    return parallel_mat_product(grouped)

# 测试用例
N = 3  # 矩阵维度
L = 5  # 矩阵数量
# 生成L个N×N随机矩阵(用GPU加速的话加.cuda())
matrices = [torch.randn(N, N).cuda() for _ in range(L)]
final_product = parallel_mat_product(matrices)
print(final_product.shape)  # 输出: torch.Size([3, 3])

方法二:批量矩阵乘法优化(效率更高)

当矩阵数量较多时,把分组后的矩阵打包成高维张量,用torch.bmm一次性完成所有组的计算,进一步提升并行效率:

import torch

def batch_parallel_mat_product(matrices):
    while len(matrices) > 1:
        last_mat = None
        # 处理奇数个矩阵的情况
        if len(matrices) % 2 != 0:
            last_mat = matrices.pop()
        
        # 把两两配对的矩阵分别打包成批量张量
        batch_left = torch.stack(matrices[::2])
        batch_right = torch.stack(matrices[1::2])
        
        # 批量计算矩阵乘积,一次完成所有组的运算
        batch_products = torch.bmm(batch_left, batch_right)
        
        # 把批量结果转回列表,准备下一轮计算
        matrices = list(batch_products)
        if last_mat is not None:
            matrices.append(last_mat)
    
    return matrices[0]

# 测试用例
N = 3
L = 6
matrices = [torch.randn(N, N).cuda() for _ in range(L)]
final_product = batch_parallel_mat_product(matrices)
print(final_product.shape)  # 输出: torch.Size([3, 3])

关键注意点

  • 并行性说明:PyTorch在GPU上执行torch.matmul或torch.bmm时,会自动利用多流并行处理不同组的计算,不需要手动额外配置多线程/进程。
  • 正确性验证:可以用逐次迭代相乘的结果和上述方法对比,确认数值一致(允许浮点误差):
# 逐次迭代计算作为对照
iter_product = matrices[0]
for mat in matrices[1:]:
    iter_product = iter_product @ mat

# 验证结果一致
print(torch.allclose(final_product, iter_product, atol=1e-6))  # 输出: True

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.13 03:56:01