Numpy切片结果不符合预期,求解析及与PyTorch的差异原因
Numpy切片维度顺序异常的原因与解决方法
异常现象
当在Numpy中同时使用基本索引和布尔高级索引进行多维数组切片时,会出现维度顺序不符合预期的情况:
import torch import numpy as np some_array = np.zeros((1, 3, 42)) chooser_mask = np.zeros((42)) # 标记要选取的2个位置 chooser_mask[13] = 1 chooser_mask[14] = 1 out_1 = some_array[0, :, chooser_mask == 1] print(out_1.shape) # 输出 (2, 3),与预期的 (3, 2) 不符
而将切片分步执行时,结果符合预期:
tmp = some_array[0] out_2 = tmp[:, chooser_mask == 1] print(out_2.shape) # 输出 (3, 2),符合预期
另外,相同逻辑在PyTorch中不会出现该问题:
some_array = torch.from_numpy(some_array) chooser_mask = torch.from_numpy(chooser_mask) out_1 = some_array[0, :, chooser_mask == 1] print(out_1.shape) # 输出 (3, 2),符合预期 tmp = some_array[0] out_2 = tmp[:, chooser_mask == 1] print(out_2.shape) # 输出 (3, 2),符合预期
原因解析
这个差异源于Numpy和PyTorch对**混合索引(基本索引+高级索引)**的处理规则不同:
Numpy规则:当切片操作中同时存在基本索引(如单个索引值
0、切片:)和高级索引(如布尔索引、整数数组索引)时,Numpy会将高级索引对应的维度优先放置在结果数组的最前面。
在some_array[0, :, chooser_mask == 1]中:0是对第一维的基本索引,会将数组从(1,3,42)降为(3,42);:是对第二维的基本索引,保留该维度;chooser_mask == 1是对第三维的布尔高级索引,Numpy会将这个索引得到的维度(长度2)前置,最终结果维度变为(2,3)。
而分步切片时,
tmp = some_array[0]已经得到(3,42)的数组,后续tmp[:, chooser_mask == 1]中只有第二维的基本索引和第三维的高级索引,此时Numpy不会调整维度顺序,因此保持(3,2)。PyTorch规则:PyTorch在处理混合索引时,会严格保留原数组的维度顺序,不会将高级索引的维度前置,因此无论是否分步切片,结果维度都符合预期。
解决方法
除了分步切片外,还可以通过以下方式避免Numpy的维度顺序异常:
- 使用索引链式调用代替一次性混合索引:
out = some_array[0][:, chooser_mask == 1] print(out.shape) # (3, 2)
- 先对目标维度进行高级索引,再去除多余维度:
out = some_array[:, :, chooser_mask == 1].squeeze(0) print(out.shape) # (3, 2)
- 使用
np.take指定索引维度:
out = np.take(some_array[0], np.where(chooser_mask == 1)[0], axis=1) print(out.shape) # (3, 2)
内容的提问来源于stack exchange,提问作者n0tis
相关产品推荐
相关产品推荐

