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

PyTorch CPU复数矩阵-向量乘法性能远逊于Numpy的问题排查

CPU上PyTorch复数矩阵-向量乘法远慢于NumPy的问题排查

现象与测试结果

在CPU环境下执行复数矩阵-向量乘法时,PyTorch的运行速度显著慢于NumPy,实/复数场景的性能对比图如下:
NumPy与PyTorch(CPU)实/复数矩阵乘法性能对比

补充测试细节

  • 多台设备均复现该现象
  • 测试过程内存无压力
  • PyTorch复数乘法未占满CPU核心(NumPy实/复数、PyTorch实数运算均能占满核心)
  • 版本信息:PyTorch 2.5.1+cu124、NumPy 1.26.4、CUDA 12.6、NVIDIA驱动560.35.03
  • 已验证二者计算结果完全一致
  • 测试均使用双精度(实部64位、复数128位),切换为单精度(torch.cfloat)仅带来小幅性能提升

测试代码

import torch
import numpy as np
import matplotlib.pyplot as plt
import time

maxn = 3000
nrep = 100

def conv(M,latype):
    if latype=='numpy':
        return np.array(M)
    if latype.startswith('torch,'):
        return torch.tensor(M,device=latype[7:])

def multtest(A,b):
    t0 = time.time()
    for i in range(nrep):
        b = A@b
    t1 = time.time()
    return (t1-t0)/nrep

ns = np.array(np.linspace(100,maxn,100),dtype=int)
numpyts = np.zeros(len(ns))
torchts = np.zeros(len(ns))

fig,axes = plt.subplots(1,2)
for ax,dtype in zip(axes,['real','complex']):
    Aorig = np.random.rand(maxn,maxn)
    borig = np.random.rand(maxn)
    if dtype == 'complex':
        Aorig = Aorig + 1.j*np.random.rand(maxn,maxn)
        borig = borig + 1.j*np.random.rand(maxn)

    for latype in ['numpy','torch, cpu']:
        A = conv(Aorig,latype)
        b = conv(borig,latype)
        ts = np.zeros(len(ns))
        for i,n in enumerate(ns):
            ts[i] = multtest(A[:n,:n],b[:n])
        ax.plot(ns,ts,label=latype)

    ax.legend()
    ax.set_title(dtype)
    ax.set_xlabel('vector/matrix size')
    ax.set_ylabel('mean matrix-vector mult time (sec)')

fig.tight_layout()
plt.show()

排查方向

  • 检查PyTorch CPU后端优化库:执行torch.__config__.show()查看当前启用的BLAS后端。NumPy通常默认绑定MKL(若系统已安装),若PyTorch使用OpenBLAS,其复数运算的多线程优化可能弱于MKL。可尝试重新编译PyTorch绑定MKL,或设置环境变量export MKL_NUM_THREADS=你的CPU核心数强制启用MKL多线程。
  • 验证多线程配置:执行torch.get_num_threads()和torch.get_num_interop_threads(),确认线程数与CPU核心数匹配。若不匹配,手动设置torch.set_num_threads(核心数)、torch.set_num_interop_threads(核心数)后重新测试。
  • 排查运算回退情况:用torch._C._debug_get_autograd_fallback_counts()查看是否有运算回退到Python实现,或通过torch.jit.trace将运算脚本化,强制启用JIT编译优化。
  • 版本兼容性测试:尝试降级或升级PyTorch版本(如2.4.x或2.6.x),确认是否为特定版本的复数运算调度Bug。
  • 系统BLAS库检查:通过np.__config__.show()查看NumPy绑定的BLAS版本,确保系统安装的MKL/OpenBLAS为最新版本,旧版本库对复数多线程运算的支持可能不足。
  • 底层BLAS性能对比:直接调用MKL的cblas_zgemv或OpenBLAS的对应复数矩阵-向量乘法函数,对比PyTorch、NumPy的封装开销,确认瓶颈是否在PyTorch的上层逻辑而非底层库。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 07:57:36