Python/Numpy脚本比MATLAB慢10倍,求加速优化建议
Python 多重循环脚本加速优化建议
问题背景
我编写了一个包含多重循环的脚本,Python版本运行耗时约88秒,而MATLAB版本仅需约9秒,目前已通过向量化移除了其中一层循环。当参数nj调整为10000时,MATLAB运行耗时21秒,Python则需248秒。相关代码如下:
MATLAB 代码
ntotal = 2000; nj = 4000; r = zeros(ntotal,nj,3); ncorr = 200; temp2 = zeros(ncorr,nj); final = zeros(3,ncorr); tic for k = 1:3 for j = 1:nj for n0 = 1:ncorr temp1=(r(1:ncorr-n0,j,k)-r(n0:ncorr-1,j,k)).^2; temp2(n0,j) = mean(temp1); end end final(k,:) = mean(temp2,2); end toc
Python 原代码
import time import numpy as np ntotal = 2000 nj = 4000 r = np.zeros((ntotal,nj,3)) ncorr = 200 temp2 = np.zeros((ncorr,nj)) final = np.zeros((3,ncorr)) t0 = time.time() for k in range(3): for j in range(nj): for n0 in range(ncorr): temp1 = (r[0:ntotal-n0,j,k]-r[n0:ntotal,j,k]) ** 2 temp2[n0,j] = np.mean(temp1) final[k,:] = np.mean(temp2,axis=1) t1 = time.time() print(t1-t0)
优化方案
1. 完全向量化运算,消除嵌套循环
把三层循环转化为numpy数组的广播操作,减少Python循环的开销:
import time import numpy as np ntotal = 2000 nj = 4000 r = np.zeros((ntotal, nj, 3)) ncorr = 200 t0 = time.time() # 调整维度顺序,让k维度优先,方便批量处理 r_reshaped = r.transpose(2, 1, 0) # shape: (3, nj, ntotal) final = np.zeros((3, ncorr)) for k in range(3): current_data = r_reshaped[k] # shape: (nj, ntotal) # 一次性计算所有n0对应的均值,再按nj维度取平均 for n0 in range(ncorr): window_len = ntotal - n0 diff_sq = (current_data[:, :window_len] - current_data[:, n0:]) ** 2 nj_mean = np.mean(diff_sq, axis=1) final[k, n0] = np.mean(nj_mean) t1 = time.time() print(t1 - t0)
2. 使用Numba JIT编译加速循环
Numba能将Python循环编译为机器码,性能接近MATLAB:
import time import numpy as np from numba import jit ntotal = 2000 nj = 4000 r = np.zeros((ntotal, nj, 3)) ncorr = 200 temp2 = np.zeros((ncorr, nj)) final = np.zeros((3, ncorr)) # 开启nopython模式,完全脱离Python解释器执行 @jit(nopython=True) def compute(r, ntotal, nj, ncorr, temp2, final): for k in range(3): for j in range(nj): for n0 in range(ncorr): temp1 = (r[:ntotal-n0, j, k] - r[n0:, j, k]) ** 2 temp2[n0, j] = np.mean(temp1) final[k, :] = np.mean(temp2, axis=1) t0 = time.time() compute(r, ntotal, nj, ncorr, temp2, final) t1 = time.time() print(t1 - t0)
3. 调整数组内存布局匹配MATLAB
MATLAB默认使用列优先(Fortran顺序)存储数组,numpy默认是行优先(C顺序),调整存储顺序可减少内存访问开销:
# 创建数组时指定order='F',和MATLAB内存布局一致 r = np.zeros((ntotal, nj, 3), order='F')
4. 利用滑动窗口工具减少重复计算
使用numpy.lib.stride_tricks.sliding_window_view生成滑动窗口,简化差值计算:
import time import numpy as np from numpy.lib.stride_tricks import sliding_window_view ntotal = 2000 nj = 4000 r = np.zeros((ntotal, nj, 3)) ncorr = 200 t0 = time.time() r_reshaped = r.transpose(2, 1, 0) final = np.zeros((3, ncorr)) for k in range(3): current_data = r_reshaped[k] for n0 in range(ncorr): # 生成两个滑动窗口,直接计算差值平方的均值 window1 = sliding_window_view(current_data, ntotal - n0, axis=1)[:, 0] window2 = sliding_window_view(current_data, ntotal - n0, axis=1)[:, n0] diff_sq = (window1 - window2) ** 2 final[k, n0] = np.mean(diff_sq) t1 = time.time() print(t1 - t0)
内容的提问来源于stack exchange,提问作者Kevin
相关产品推荐
相关产品推荐

