如何用Scipy RegularGridInterpolator沿指定轴插值多维数组?
多维数据立方体沿指定维度插值(保持其余维度不变)
我正尝试对尺寸为NxNyNz的数据立方体执行简单线性插值,要求保持另外两个维度不变,得到尺寸为NxnewNyNz的新数据立方体。我了解到Scipy的RegularGridInterpolator是合适工具,但不清楚如何生成输入参数。
以下是最小可复现代码(MWE):
from scipy.interpolate import RegularGridInterpolator import numpy as np x = np.linspace(1,4,11) y = np.linspace(4,7,22) z = np.linspace(7,9,33) V = np.zeros((11,22,33)) for i in range(11): for j in range(22): for k in range(33): V[i,j,k] = 100*x[i] + 10*y[j] + z[k] fn = RegularGridInterpolator((x,y,z), V) pts = np.array([[[[2,6,8],[3,5,7]], [[2,6,8],[3,5,7]]]]) out = fn(pts) print(out, out.shape)在该示例中,我希望使用新点
xnew = np.linspace(2,3,50),同时保持y和z不变,使输出数组形状为(50,22,33)。另外,如何将该方法推广到n维数组沿某一维插值,同时保持其余坐标不变?
针对3维数据的解决方案
要实现沿x轴插值并保持y、z维度不变,核心是生成所有(xnew_i, y_j, z_k)的坐标组合,通过插值函数计算后重塑回目标形状:
from scipy.interpolate import RegularGridInterpolator import numpy as np # 原始数据定义 x = np.linspace(1,4,11) y = np.linspace(4,7,22) z = np.linspace(7,9,33) V = np.zeros((11,22,33)) for i in range(11): for j in range(22): for k in range(33): V[i,j,k] = 100*x[i] + 10*y[j] + z[k] # 初始化插值器 interp_fn = RegularGridInterpolator((x, y, z), V) # 定义新的x轴坐标 xnew = np.linspace(2, 3, 50) # 生成插值所需的全量坐标网格 Xnew, Y, Z = np.meshgrid(xnew, y, z, indexing='ij') # 将网格展平并组合成(n_points, 3)的点数组,适配插值器输入格式 pts = np.stack([Xnew.ravel(), Y.ravel(), Z.ravel()], axis=1) # 执行插值并重塑为目标形状 Vnew = interp_fn(pts).reshape(xnew.size, y.size, z.size) print(Vnew.shape) # 输出 (50, 22, 33)
代码说明
np.meshgrid(..., indexing='ij'):确保生成的网格维度顺序与原始数据一致(x在前,y、z在后)np.stack(..., axis=1):将展平后的各维度坐标组合成插值器要求的(n_points, 维度数)格式reshape:把一维插值结果恢复为目标三维结构
推广到n维数组的通用方法
可以编写一个通用函数,支持任意维度数组沿指定轴插值,自动保留其余维度的原始结构:
def interpolate_along_axis(original_coords, data, axis, new_coords): """ 沿指定维度插值,保持其余维度不变 参数: original_coords: 元组,每个元素对应数据各维度的原始坐标数组 data: n维numpy数组,原始数据 axis: int,要插值的维度索引(从0开始计数) new_coords: 1维numpy数组,该维度的新坐标序列 返回: n维numpy数组,插值后的数据,指定维度长度为new_coords.size """ from scipy.interpolate import RegularGridInterpolator # 初始化插值器 interp_fn = RegularGridInterpolator(original_coords, data) # 构建新的网格组件:替换指定维度为new_coords,其余保留原始坐标 grid_components = list(original_coords) grid_components[axis] = new_coords # 生成全量网格,保证维度顺序与原始数据一致 grid = np.meshgrid(*grid_components, indexing='ij') # 展平网格并组合成插值器所需的点数组 pts = np.stack([g.ravel() for g in grid], axis=1) # 计算插值并重塑为目标形状 new_shape = list(data.shape) new_shape[axis] = new_coords.size interpolated_data = interp_fn(pts).reshape(new_shape) return interpolated_data
通用函数使用示例
# 沿第0维(x轴)插值,得到(50,22,33)的结果 Vnew_3d = interpolate_along_axis((x, y, z), V, axis=0, new_coords=xnew) print(Vnew_3d.shape) # (50, 22, 33) # 测试沿y轴插值,生成(11,100,33)的数组 ynew = np.linspace(4.5, 6.5, 100) Vnew_y = interpolate_along_axis((x, y, z), V, axis=1, new_coords=ynew) print(Vnew_y.shape) # (11, 100, 33)
内容的提问来源于stack exchange,提问作者Georges Leukic
相关产品推荐
相关产品推荐

