如何使用数组索引对3D NumPy数组按指定维度切片并得到目标结果
解决方法
你要实现的是对3D数组第0维的每个切片,使用独立的行、列索引范围进行切片,numpy原生的冒号切片不支持传入数组作为起止值,可以用以下两种方案实现:
方案1:广播花式索引(高效无循环,适合大数据量)
利用numpy的广播机制构造三维索引数组,一次性取出所有需要的元素:
import numpy as np ts = np.arange(25*3).reshape(3,5,5) newr1 = np.array([1,0,2]) newr2 = np.array([3,2,4]) newc1 = np.array([1,2,0]) newc2 = np.array([3,4,2]) # 适配numpy切片左闭右开规则,右边界需要+1才能匹配你给出的目标输出 adj_newr2 = newr2 + 1 adj_newc2 = newc2 + 1 # 构造各维度索引 dim0_idx = np.arange(ts.shape[0])[:, None, None] # 形状(3,1,1),对应第0维的3个元素 # 构造行索引:每个第0维元素对应独立的行范围,形状(3, 行长度, 1) row_len = (adj_newr2 - newr1)[0] # 这里所有行范围长度一致,均为3 row_idx = (newr1[:, None] + np.arange(row_len))[..., None] # 构造列索引:每个第0维元素对应独立的列范围,形状(3, 1, 列长度) col_len = (adj_newc2 - newc1)[0] # 这里所有列范围长度一致,均为3 col_idx = (newc1[:, None] + np.arange(col_len))[:, None, :] # 索引取值 result = ts[dim0_idx, row_idx, col_idx] print(result)
运行后输出和你给出的目标结果完全一致。如果你的行/列范围长度每个第0维元素不同,可以用下面的列表推导式方案。
方案2:列表推导式(写法简单,适合小数据量)
直接遍历第0维的每个元素,分别切片后堆叠:
result = np.stack([ ts[i, r1:r2+1, c1:c2+1] for i, (r1, r2, c1, c2) in enumerate(zip(newr1, newr2, newc1, newc2)) ])
写法简洁易懂,数据量不大的时候性能完全够用。
内容的提问来源于stack exchange,提问作者user1946217
相关产品推荐
相关产品推荐

