如何沿指定轴对1D数组与多维数组进行并行线性插值?
向量化实现沿特定轴的多维数组插值
可以通过scipy.interpolate.interp1d或纯NumPy的向量化操作来避免双重循环,优雅实现需求。以下是两种可行方案:
方案一:使用Scipy的interp1d(简洁直观)
interp1d支持多维数组,只需指定插值轴即可批量处理所有维度的插值任务:
import numpy as np from scipy.interpolate import interp1d def foo_vectorized(T, B, x_in): len_T = len(T) len_x_in = len(x_in) # 构造x和y数组(与原代码逻辑一致) x = np.ones((6, len_T, len_x_in)) x[0, :, :] = 0 x[1, :, :] = 0.3 - 0.1*T[:, np.newaxis] x[2, :, :] = 0.35 - 0.1*T[:, np.newaxis] x[3, :, :] = 0.8 - 0.2*T[:, np.newaxis] x[4, :, :] = 0.9 - 0.2*T[:, np.newaxis] x[5, :, :] = 1.0 y = np.ones((6, len_T, len_x_in)) y[0, :, :] = -0.25*T[:, np.newaxis]*(1 + B) y[1, :, :] = -1 y[2, :, :] = 1 y[3, :, :] = 1 y[4, :, :] = -1 y[5, :, :] = -1 # 创建插值函数,指定沿第0轴(原问题的第一轴)插值 interpolator = interp1d(x, y, axis=0) # 构造插值点网格:每个(i,j)位置对应x_in[j] x_in_grid = np.broadcast_to(x_in, (len_T, len_x_in)) # 执行插值,直接得到(2,10)形状的结果 return interpolator(x_in_grid)
方案二:纯NumPy向量化实现(无Scipy依赖)
利用np.searchsorted定位插值位置,结合广播机制手动计算线性插值:
import numpy as np def foo_numpy_vectorized(T, B, x_in): len_T = len(T) len_x_in = len(x_in) # 构造x和y数组(与原代码逻辑一致) x = np.ones((6, len_T, len_x_in)) x[0, :, :] = 0 x[1, :, :] = 0.3 - 0.1*T[:, np.newaxis] x[2, :, :] = 0.35 - 0.1*T[:, np.newaxis] x[3, :, :] = 0.8 - 0.2*T[:, np.newaxis] x[4, :, :] = 0.9 - 0.2*T[:, np.newaxis] x[5, :, :] = 1.0 y = np.ones((6, len_T, len_x_in)) y[0, :, :] = -0.25*T[:, np.newaxis]*(1 + B) y[1, :, :] = -1 y[2, :, :] = 1 y[3, :, :] = 1 y[4, :, :] = -1 y[5, :, :] = -1 # 扩展x_in形状以匹配广播要求 x_in_expanded = x_in[np.newaxis, np.newaxis, :] # 找到每个插值点在x轴上的插入位置(沿第0轴) indices = np.searchsorted(x, x_in_expanded, side='right', axis=0) # 处理边界情况:限制索引在合法范围内 indices = np.clip(indices, 1, 5) idx_left, idx_right = indices - 1, indices # 创建网格索引,用于批量获取左右邻点的x/y值 i_grid = np.arange(len_T)[np.newaxis, :, np.newaxis] j_grid = np.arange(len_x_in)[np.newaxis, np.newaxis, :] x_left = x[idx_left, i_grid, j_grid].squeeze() x_right = x[idx_right, i_grid, j_grid].squeeze() y_left = y[idx_left, i_grid, j_grid].squeeze() y_right = y[idx_right, i_grid, j_grid].squeeze() # 计算线性插值结果 slope = (y_right - y_left) / (x_right - x_left) return y_left + slope * (x_in - x_left)
两种方案均输出(2,10)形状的结果,且与原循环版本的计算结果一致(浮点误差范围内),性能远优于双重循环实现。
内容的提问来源于stack exchange,提问作者GioR
相关产品推荐
相关产品推荐

