使用Torch张量作为NumPy高级索引时维度被压缩的原因探究
NumPy与PyTorch索引行为差异分析
测试代码
import numpy as np import torch a_numpy = np.array([1,2,3]) a_torch = torch.tensor([1,2,3]) index_torch = torch.tensor([[0]]) index_np = np.array([[0]]) print(f"{a_numpy[index_torch]=}") print(f"{a_numpy[index_np]=}") print(f"{a_torch[index_torch]=}") print(f"{a_torch[index_np]=}")
执行输出
a_numpy[index_torch]=np.int64(1) a_numpy[index_np]=array([[1]]) a_torch[index_torch]=tensor([[1]]) a_torch[index_np]=tensor([[1]])
原因解析
这不是BUG,是NumPy和PyTorch对非原生索引类型的处理逻辑差异导致的:
- NumPy用原生数组索引:当用
index_np(NumPy二维数组)索引时,严格遵循多维数组索引规则,返回和索引维度一致的二维数组结果array([[1]])。 - NumPy用PyTorch张量索引:NumPy不直接支持PyTorch张量作为索引,会先把张量转换成Python嵌套列表
[[0]]。此时NumPy将这个嵌套列表视为多个标量索引的组合,而非二维索引,最终返回单个标量np.int64(1)。 - PyTorch的统一处理:不管索引是PyTorch张量还是NumPy数组,PyTorch都会先将其转换为张量,然后按照自身的多维索引规则处理,所以两种索引方式都返回二维张量
tensor([[1]])。
简单来说,这是NumPy对非原生索引类型的转换逻辑导致的行为差异,属于设计层面的既定规则,而非程序BUG。
内容的提问来源于stack exchange,提问作者big_noob
相关产品推荐
相关产品推荐

