PyTorch张量切片疑问:a[b]为何返回[B,3,3]形状?
解释PyTorch中
a[b]的索引行为 这是PyTorch中**高级索引(Advanced Indexing)**的典型表现,和你熟悉的常规切片逻辑完全不同,具体原理拆解如下:
1. 原张量与索引张量的维度对应
- 原张量
a的形状是[A, 3]:第0维是长度为A的样本轴,第1维是每个样本的3维特征轴。 - 索引张量
b的形状是[B, 3]且类型为long:这意味着b中的每一个元素,都是用来索引a的第0维(样本轴)的位置值。
2. 索引的执行逻辑
当执行a[b]时,PyTorch会逐元素处理索引张量b:
- 对于
b中每一个位置的元素b[i][j](i从0到B-1,j从0到2),PyTorch会取出a中第0维位置为b[i][j]的整个样本,也就是形状为[3]的张量。 - 这些取出的
[3]张量会按照b的形状[B, 3]排列,同时保留原张量a的剩余维度(特征轴的长度3),最终拼接成形状为[B, 3, 3]的结果。
举个具体小例子:
假设a = torch.rand(5,3)(A=5),b = torch.tensor([[0,2,4], [1,3,0]])(B=2),那么:
a[b[0]]会取出a[0], a[2], a[4],得到形状[3,3]的张量;a[b[1]]会取出a[1], a[3], a[0],同样得到形状[3,3]的张量;- 最终
a[b]就是把这两个[3,3]的张量堆叠起来,形成[2,3,3]的结果,和你示例中的形状逻辑一致。
3. 和常规切片的核心区别
常规切片(比如a[N1:N2:jump, [0,2]])是对原张量的维度直接做范围/间隔选取,结果的维度数和原张量一致;而高级索引是用张量作为索引,将原张量的元素按索引张量的形状重新排列,结果的维度数等于索引张量的维度数加上原张量被索引维度之外的剩余维度数。
注意:如果b中的元素超出a第0维的合法范围(即小于0或大于等于A),PyTorch会抛出索引越界错误。你的示例中因为torch.rand(20,3)生成的是0到1之间的浮点数,转long后全部为0,所以不会触发错误。
内容的提问来源于stack exchange,提问作者Mohit Lamba
相关产品推荐
相关产品推荐

