Python中使用另一多维数组对三维NumPy数组按深度维度采样生成二维数组的最简方法
解决NumPy三维数组按位置随机采样深度维度的问题
嘿,这个坑我之前也踩过!你遇到的核心问题是NumPy高级索引的维度匹配规则——直接用A[np.random.randint(0, 39, (23,23))]的话,NumPy会把这个二维索引数组当成第一维(高度H)的索引,导致结果完全不符合预期,甚至会生成一个四维数组。
下面给你两种最简可行的方案,都是经过验证的:
方案一:用np.take_along_axis(推荐,直观简洁)
这个函数专门设计用来沿指定轴提取元素,完美适配你这种“每个前两维位置对应一个第三维索引”的场景:
import numpy as np # 定义原三维数组(H=23, W=23, D=39) A = np.random.randint(0, 10, (23, 23, 39)) # 生成每个(H,W)位置对应的深度维度随机索引,形状(23,23) depth_idx = np.random.randint(0, 39, size=A.shape[:2]) # 给索引数组增加一个维度,使其和原数组的第三维匹配(变成(23,23,1)) # 用take_along_axis沿第2轴(深度轴)提取元素,最后去掉多余的维度得到二维数组 B = np.take_along_axis(A, depth_idx[..., np.newaxis], axis=2).squeeze() # 验证结果形状:输出(23, 23) print(B.shape)
如果不想手动加维度,也可以生成索引时直接用keepdims=True:
depth_idx = np.random.randint(0, 39, size=A.shape[:2], keepdims=True) B = np.take_along_axis(A, depth_idx, axis=2).squeeze()
方案二:用元组式高级索引(手动对齐维度)
如果你更习惯手动处理索引维度,可以用这种方式,原理是给前两维也生成对应形状的索引数组,和深度索引一起组成元组来索引:
import numpy as np A = np.random.randint(0, 10, (23, 23, 39)) depth_idx = np.random.randint(0, 39, (23,23)) # 生成高度维度的索引:形状(23,1),会和宽度维度自动广播 h_idx = np.arange(A.shape[0])[:, None] # 生成宽度维度的索引:形状(23,) w_idx = np.arange(A.shape[1]) # 用元组索引,每个(h,w)位置取depth_idx[h,w]对应的深度元素 B = A[h_idx, w_idx, depth_idx] # 验证结果形状:输出(23, 23) print(B.shape)
为什么原来的方法不行?
简单说,A[idx]这种写法中,NumPy会默认把idx当成**第一维(H)**的索引,相当于对每个idx[h,w],都取A[idx[h,w], :, :]这个二维数组,最后拼接成一个(23,23,23,39)的四维数组,完全不是你想要的结果。而上面两种方法都是明确指定了每个位置对应第三维的索引,完美对齐了维度。
内容的提问来源于stack exchange,提问作者user8208938
相关产品推荐
相关产品推荐

