如何优化使用scipy样条插值的numpy高耗时循环代码
NumPy循环代码优化方案
你的原代码耗时高的核心原因是存在大量无意义的重复计算,我们可以通过优化完全消除循环,把运行时长压缩到0.1秒以内:
优化核心逻辑
- 你的
x、x1都是按列复制的完全相同的数组,tck0、tck1是预先计算好的固定样条参数,因此两个插值结果仅需要计算1次即可,不需要在10000次循环里重复调用splev - 阈值判断和加权求和的操作可以直接利用NumPy的广播特性全量计算,不需要逐列处理
优化后代码
import numpy as np from scipy import interpolate import numpy.matlib as matlib I=10000 T=10000 y = np.random.uniform(0, 10, (I,T)) # x每列都相同,不需要重复生成 x_col = np.linspace(0,25,I) # x1每列都相同,不需要重复生成 x1_col = np.linspace(2,20,I) # 预计算两个固定样条 tck0 = interpolate.splrep(x_col, y[:,0], s=0) tck1 = interpolate.splrep(x_col, y[:,1], s=0) # 仅计算一次插值结果,reshape为列向量后广播匹配y的形状 splev0 = interpolate.splev(x1_col, tck0, der=0, ext=0).reshape(-1,1) splev1 = interpolate.splev(x1_col, tck1, der=0, ext=0).reshape(-1,1) # 全量向量化计算,完全消除循环 ct = (y <=5)*splev0 + (y>5)*splev1
额外优化建议
如果你的实际场景中tck不是固定的(比如每一列都需要单独拟合样条),可以考虑用scipy.interpolate.interp1d搭配向量化调用,或者用Numba装饰循环进一步提速。
内容的提问来源于stack exchange,提问作者mararipe
相关产品推荐
相关产品推荐

