You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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] = 5
  • b[0][1] = 1 → 行0、列1 → a[0][1] = 5
  • b[1][0] = 1 → 行1、列1 → a[1][1] = 8
  • b[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] = 7
  • b[0][1] = 1 → 行1、列1 → a[1][1] = 8
  • b[1][0] = 1 → 行1、列0 → a[1][0] = 7
  • b[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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.29 00:54:29