使用NumPy数组索引PyTorch Tensor是否属于合法操作?
结论
该现象是PyTorch 1.9.0、1.9.1两个版本的独有bug,并非设计层面的非预期行为。PyTorch原生支持NumPy数组作为Tensor的索引输入,正常版本下a[b.numpy()]应当和a[b]的执行逻辑、结果完全一致。
根因说明
1.9.x版本对NumPy类型索引的维度校验逻辑存在缺陷:当索引数组的维度参数命中特定阈值时,会错误地将单个数组索引拆分为多个独立的维度索引序列,最终抛出「索引数量超过Tensor维度」的报错。你测试的可运行/不可运行维度组合,刚好卡在该bug的判断边界上。
验证方案
将PyTorch版本升级至1.10及以上,或降级至1.8.x及以下,再运行你提供的测试代码,所有维度的用例均可正常执行,不会触发该错误。
临时规避方法
如果暂时无法升级PyTorch版本,手动将NumPy索引转换为Tensor后再传入即可避开该bug:
a[torch.from_numpy(b.numpy())]
内容的提问来源于stack exchange,提问作者Proko
相关产品推荐
相关产品推荐

