PyTorch中不同嵌套列表与张量的索引机制解析
PyTorch索引两种结果的底层逻辑解析
先明确原始张量:
import torch x = torch.arange(16).reshape(4, 4) # 输出: # tensor([[ 0, 1, 2, 3], # [ 4, 5, 6, 7], # [ 8, 9, 10, 11], # [12, 13, 14, 15]])
你的例子里出现了两种完全不同的索引结果,核心差异在于索引的类型(Python列表 vs PyTorch张量)以及嵌套列表的深度对应的索引维度解析规则,下面分情况拆解:
情况1:得到单个元素组成的张量(tensor([2, 7]))
对应这两种索引方式:
# 两个独立列表作为索引参数 x[[0,1], [2,3]] # 二维嵌套列表作为单个索引参数 x[[[0,1], [2,3]]]
底层逻辑:
这本质是PyTorch的高级索引中的“配对索引”,规则如下:
- 当传入多个等长的一维列表(如
[0,1]和[2,3]),PyTorch会将列表中对应位置的索引配对,取出x[行索引[i], 列索引[i]]的元素:- 第0对:
x[0,2] = 2 - 第1对:
x[1,3] = 7
最终拼接成形状为(2,)的张量。
- 第0对:
- 当传入二维嵌套列表
[[0,1],[2,3]]作为单个索引参数时,PyTorch会自动将其解析为两个维度的索引组合(外层列表的两个子列表分别对应行、列维度的索引),行为和传入两个独立列表完全一致,因此得到相同结果。
情况2:得到子张量块(形状为(2,2,4)的张量)
对应这两种索引方式:
# 三维嵌套列表作为单个索引参数 x[[[[0,1], [2,3]]]] # 二维PyTorch张量作为索引参数 x[torch.tensor([[0,1], [2,3]])]
底层逻辑:
这是PyTorch的维度选择索引,规则如下:
- 三维嵌套列表的情况:当嵌套深度超过原张量的维度数(原张量是2维)时,PyTorch会将索引的每个元素看作是原张量第一个维度的行索引,并保持索引的嵌套结构:
- 索引
[[[0,1], [2,3]]]的结构对应两组行索引:第一组[0,1]、第二组[2,3]。 - 取出每组对应的行张量,按索引结构堆叠,最终得到形状为
(2,2,4)的子张量块。
- 索引
- PyTorch张量作为索引的情况:张量索引的规则和列表索引完全不同——张量的每个元素都是原张量第一个维度的索引值,最终结果的形状是索引张量的形状 + 原张量剩余维度的形状:
- 索引张量
torch.tensor([[0,1], [2,3]])的形状是(2,2),原张量剩余维度是(4,),所以最终结果形状是(2,2,4)。 - 具体来说,就是取出
x[0]、x[1]、x[2]、x[3],然后按照索引张量的(2,2)结构排列,每个位置放对应的行张量,就得到了例子中的子张量块。
- 索引张量
关键总结
- Python列表索引:
- 当嵌套深度等于原张量维度数时,会被解析为多维度的配对索引,取出单个元素。
- 当嵌套深度超过原张量维度数时,会被解析为维度上的分组选择,取出子张量块。
- PyTorch张量索引:
- 无论张量维度多少,每个元素都是原张量第一个维度的索引,结果形状是「索引形状 + 原张量剩余维度形状」,取出对应行/子张量的堆叠。
内容的提问来源于stack exchange,提问作者Lovkush
相关产品推荐
相关产品推荐

