为何洛伦兹函数向量化实现lorz1/lorz2比循环版lorz3更慢?
我编写了一个根据freq、fwhm、amp计算洛伦兹函数(Lorentzian)的函数,希望将其向量化,以实现对一组freq、fwhm和amp的批量计算:
import numpy as np def lorz1(freq_series, freq, fwhm, amp): numerator = fwhm denominator = (2*np.pi) * ((freq_series[:,None] - freq)**2 + fwhm**2/4) lor = numerator / denominator main_peak = amp*(lor/np.linalg.norm(lor, axis=0)) return np.sum(main_peak, axis=1) def lorz2(freq_series, freq, fwhm, amp): numerator = fwhm[:,None] denominator = (2*np.pi) * ((freq_series - freq[:,None])**2 + fwhm[:,None]**2/4) lor = numerator / denominator main_peak = amp[:,None]*(lor/np.linalg.norm(lor, axis=1)[:,None]) return np.sum(main_peak, axis=0) def lorz3(freq_series, freq, fwhm, amp): numerator = fwhm denominator = (2*np.pi) * ((freq_series - freq)**2 + fwhm**2/4) lor = numerator / denominator main_peak = amp*(lor/np.linalg.norm(lor)) return main_peak series = np.linspace(0,100,50000) freq = np.random.uniform(5,50,50) fwhm = np.random.uniform(0.01,0.05,50) amps = np.random.uniform(5,500,50)
计时结果
%timeit lorz1(series, freq, fwhm, amps)
每次循环耗时38.4 ms ± 1.7 ms(7次运行,每次10个循环的均值±标准差)
%timeit lorz2(series, freq, fwhm, amps)
每次循环耗时29.8 ms ± 1.8 ms(7次运行,每次10个循环的均值±标准差)
%timeit np.sum(np.array([lorz3(series, item1, item2, item3) for (item1,item2,item3) in zip(freq, fwhm, amps)]), axis=0)
每次循环耗时24.1 ms ± 5.02 ms(7次运行,每次10个循环的均值±标准差)
问题
我在lorz1和lorz2的向量化实现中哪里出错了?它们不是应该比lorz3更快吗?
原因分析与优化建议
1. 内存开销是核心问题
lorz1和lorz2通过广播生成了50000×50的超大中间数组(series长度50000,待处理的峰数量50),这类数组仅float64类型就占用约20MB内存,且计算过程中会不断生成同尺寸的临时数组(平方、加法、除法等步骤)。内存读写的开销远超过向量化带来的计算效率提升,反而拖慢了整体速度。
而lorz3的列表推导式是逐个处理每个峰,每次仅生成50000长度的小型数组,内存占用低,CPU缓存命中率更高,实际运行效率更好。
2. 归一化步骤的额外开销
lorz1和lorz2的归一化操作(np.linalg.norm+广播除法)需要对大数组做全局计算和广播,会额外生成大尺寸临时数组,进一步增加内存压力和计算耗时。而lorz3仅对单个峰的小数组做归一化,计算量分散,开销更低。
3. 优化方案
如果想兼顾向量化的简洁性和性能,可以采用逐个生成峰并累加的方式,既利用numpy对单个峰的向量化计算,又避免大数组的内存开销:
def lorz_opt(freq_series, freq, fwhm, amp): total = np.zeros_like(freq_series) for f, w, a in zip(freq, fwhm, amp): numerator = w denominator = (2*np.pi) * ((freq_series - f)**2 + (w**2)/4) lor = numerator / denominator total += a * (lor / np.linalg.norm(lor)) return total
这个版本的性能和lorz3的列表推导式相当,甚至略优(避免了列表转数组的额外开销)。
内容的提问来源于stack exchange,提问作者Prasad Mani

