PyTorch张量索引疑问:为何不能用':'替代torch.arange?
问题解析:PyTorch批量索引的维度对齐逻辑
先看你的示例数据:
import torch a = torch.arange(48).reshape((3,4,4)) coords = torch.tensor([[0,1],[1,2],[1,3]],dtype=int)
为什么a[:,coords[:,0],coords[:,1]]不符合预期?
当你用:作为第一个维度的索引时,相当于把第一个维度的所有3个batch都取出来,此时第二个维度的coords[:,0]是形状为[3]的张量,第三个维度的coords[:,1]也是[3]的张量。
PyTorch的索引规则中,当多个维度的索引张量形状不匹配时,会触发广播机制:它会把这两个[3]的张量各自扩展成[3,3]的形状。这意味着,每个batch都会取coords[:,0]里的所有3个行索引,和coords[:,1]里的所有3个列索引的组合,最终得到的是一个[3,3]的张量,而不是你想要的每个batch取一个元素的[3]张量。
正确写法a[torch.arange(3),coords[:,0],coords[:,1]]的逻辑
这里第一个维度用torch.arange(3)生成了一个形状为[3]的索引张量,和后面两个维度的索引张量coords[:,0]、coords[:,1](都是[3]形状)完全对齐。
此时PyTorch会按位置一一对应取元素:
- 第0个batch,取
(coords[0,0], coords[0,1])即(0,1)位置的元素 - 第1个batch,取
(coords[1,0], coords[1,1])即(1,2)位置的元素 - 第2个batch,取
(coords[2,0], coords[2,1])即(1,3)位置的元素
最终得到形状为[3]的张量,正好是每个batch对应一个目标元素的结果。
内容的提问来源于stack exchange,提问作者mt-clemente
相关产品推荐
相关产品推荐

