PyTorch CPU复数矩阵-向量乘法性能远逊于Numpy的问题排查
CPU上PyTorch复数矩阵-向量乘法远慢于NumPy的问题排查
现象与测试结果
在CPU环境下执行复数矩阵-向量乘法时,PyTorch的运行速度显著慢于NumPy,实/复数场景的性能对比图如下:
补充测试细节
- 多台设备均复现该现象
- 测试过程内存无压力
- 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
相关产品推荐
相关产品推荐

