torch.gather()如何构造新张量?函数运行原理解析
torch.gather 运行逻辑通俗解释 核心本质:沿着指定维度,按索引张量给定的位置从原张量中提取元素,输出张量的形状和索引张量完全一致。
先明确示例用到的两个张量结构:
a = torch.tensor([[4,5,6],[7,8,9],[10,11,12]]) # a的3行3列结构,dim=0对应行方向,dim=1对应列方向: # 行0: [4, 5, 6] # 行1: [7, 8, 9] # 行2: [10, 11, 12] b = torch.tensor([[1,1],[1,2]]) # b是2行2列的索引张量,和最终输出形状相同
dim=1 时的计算规则
当指定dim=1,代表沿列方向取值:
- 对于索引张量
b中任意位置(i,j)的元素值idx = b[i][j] - 行坐标保持和索引位置的行号
i一致,列坐标替换为idx - 输出对应位置的值为
a[i][idx]
逐位置计算结果:
b[0][0] = 1→ 行0、列1 →a[0][1] = 5b[0][1] = 1→ 行0、列1 →a[0][1] = 5b[1][0] = 1→ 行1、列1 →a[1][1] = 8b[1][1] = 2→ 行1、列2 →a[1][2] = 9
最终输出和运行结果一致:
tensor([[5, 5], [8, 9]])
dim=0 时的计算规则
当指定dim=0,代表沿行方向取值:
- 对于索引张量
b中任意位置(i,j)的元素值idx = b[i][j] - 列坐标保持和索引位置的列号
j一致,行坐标替换为idx - 输出对应位置的值为
a[idx][j]
逐位置计算结果:
b[0][0] = 1→ 行1、列0 →a[1][0] = 7b[0][1] = 1→ 行1、列1 →a[1][1] = 8b[1][0] = 1→ 行1、列0 →a[1][0] = 7b[1][1] = 2→ 行2、列1 →a[2][1] = 11
最终输出和运行结果一致:
tensor([[ 7, 8], [ 7, 11]])
通用记忆方法
不管是几维张量、指定哪个dim,都遵循三个固定规则:
- 输出张量的形状永远和传入的index索引张量形状完全相同
- 除了指定的dim维度,其余所有维度的坐标,都和索引张量中对应元素的位置坐标保持一致
- 指定dim维度的坐标值,直接替换为索引张量中存储的数值
比如三维张量下指定dim=2,索引张量中位置(i,j,k)存储的值为idx,则输出对应位置的值就是a[i][j][idx]。
内容的提问来源于stack exchange,提问作者James Arten
相关产品推荐
相关产品推荐

