如何通过3D索引实现NumPy 4D数组到3D数组的广播?
问题描述
我有维度为(time, level, lat, lon)的4D矩阵,想基于level轴(axis=1)的最大值位置,将其降维为(time, lat, lon)的3D矩阵。目前通过三重循环得到了正确结果,但尝试广播索引时要么报错要么维度不对,求无循环的高效实现方式。
简化示例代码:
#!/usr/bin/env ipython import numpy as np # ---------------------- nx = 10 ny = 10 ntime = 20 nlevs = 30 # ============================================= np.random.seed(10); datain_a = np.random.random((ntime,nlevs,ny,nx)); # generate data A datain_b = np.random.random((ntime,nlevs,ny,nx)); # generate data B # --------------------------------------------- calc_smt = np.abs(np.diff(datain_a,axis=1)/np.diff(datain_b,axis=1)) # calculate some ratio between two matrices # --------------------------------------------- calc_a = np.nanmax(calc_smt,axis=1) # find the maximum ratio at every gridcell -- answer has dimensions (time,lat,lon) ind_out = np.argmax(calc_smt,axis=1) # location of maximum ratio # --------------------------------------------------------------------------------- # Broadcasting attempts: calc_b = datain_b[ind_out] # Get an error with axis 26 is out of bounds... # NOT WORKING calc_b = datain_b[ind_out[:,np.newaxis,:,:]] # still an error: IndexError: index 26 is out of bounds for axis 0 with size 20 # NOT WORKING # --------------------------------------------------------------------------- # let us try 1st time moment: dd_a = ind_out[0,:,:] dd_b = datain_b[0,:,:,:] smt = dd_b[dd_a] # getting something with dimensions (10,10,10,10) # NOT WORKING? # --------------------------------------------------------------------------- # This is the output I expect: correct_output = np.zeros((ntime,ny,nx)); for itime in range(ntime): for jj in range(ny): for ii in range(nx): correct_output[itime,jj,ii] = datain_b[itime,ind_out[itime,jj,ii],jj,ii] # ------------------------------------------------------------------------------ # How to get the same without 3 loops?
高效无循环实现方法
方法1:使用np.take_along_axis(推荐)
这是NumPy专门为这类"按索引选取对应轴元素"场景设计的函数,语法简洁且高效:
# 给ind_out添加对应level轴的维度,使其和datain_b的维度匹配 ind_expanded = ind_out[:, np.newaxis, :, :] # 沿着axis=1(level轴)选取对应索引的元素 calc_b = np.take_along_axis(datain_b, ind_expanded, axis=1) # 去掉多余的level维度,得到(time, lat, lon)的结果 calc_b = calc_b.squeeze(axis=1)
验证结果一致性:
print(np.allclose(calc_b, correct_output)) # 输出True
方法2:手动构造全维度索引
如果不想用take_along_axis,可以手动为每个轴构造索引数组,利用广播机制实现选取:
# 构造time轴的索引:(ntime, 1, 1),和ind_out广播匹配 time_idx = np.arange(ntime)[:, np.newaxis, np.newaxis] # 构造lat和lon轴的索引:(1, ny, nx) lat_idx = np.arange(ny)[np.newaxis, :, np.newaxis] lon_idx = np.arange(nx)[np.newaxis, np.newaxis, :] # 直接索引得到目标结果 calc_b = datain_b[time_idx, ind_out, lat_idx, lon_idx]
单时间步的正确处理
针对你测试的首个时间步场景,正确的索引方式如下:
dd_a = ind_out[0,:,:] dd_b = datain_b[0,:,:,:] # 构造lat和lon的索引数组 lat_idx = np.arange(ny)[:, np.newaxis] lon_idx = np.arange(nx)[np.newaxis, :] # 正确选取对应level的元素 smt = dd_b[dd_a, lat_idx, lon_idx] print(smt.shape) # 输出(10,10),符合预期
内容的提问来源于stack exchange,提问作者msi_gerva
相关产品推荐
相关产品推荐

