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

PyTorch中基于同形状张量Argmax结果的张量索引问题

在PyTorch中用Argmax结果索引同形状张量的正确方法

这个问题属于PyTorch高级索引里的典型场景,我来帮你解决并解释清楚其中的逻辑:

核心需求回顾

你有两个形状为[3,4]的张量x和y:

  • 从x的每行中取出最大值的列索引(x_argmax,形状[3],值为[2,0,1])
  • 要从y的对应行中,取出这些索引位置的元素,最终得到[2,4,9]

正确解法1:组合高级索引

直接构造匹配的行索引和列索引,一起传入张量索引:

import torch as T

x = T.tensor([[1, 2, 8, 3], [6, 3, 3, 5], [2, 8, 1, 7]])
y = T.tensor([[0, 1, 2, 3], [4, 5, 6, 7], [8, 9, 10, 11]])

x_max, x_argmax = T.max(x, dim=1)

# 构造行索引:0,1,2 对应y的每一行
row_indices = T.arange(y.shape[0])
# 组合行和列索引,取出(0,2), (1,0), (2,1)位置的元素
result = y[row_indices, x_argmax]

print(result)  # 输出:tensor([2, 4, 9])

为什么之前的尝试不对?

我来逐个分析你试的方法:

  • y[x_argmax]:把x_argmax的元素当作行索引,取出的是y的第2、0、1行,得到的是[3,4]的张量,不是每行对应列的单个元素
  • y[:, x_argmax]:对每一行都取x_argmax指定的3列,得到的是[3,3]的张量,不符合需求
  • y[..., x_argmax]:和上面的[:, x_argmax]效果完全一致,...在这里等价于所有行的切片
  • y[x_argmax.unsqueeze(1)]:把x_argmax变成[3,1]后,依然是取第2、0、1行,结果是[3,1,4]的张量,不是目标值

正确解法2:使用torch.gather

PyTorch提供了专门的gather函数,用于按索引在指定维度收集元素:

# index需要和y的维度数一致,所以先把x_argmax变成[3,1]
result = y.gather(dim=1, index=x_argmax.unsqueeze(1)).squeeze()
print(result)  # 输出:tensor([2, 4, 9])

gather的逻辑解释:

  • dim=1表示我们要在列维度上收集元素
  • index参数的形状需要和原张量的维度数匹配(这里y是2D,所以index也要是2D),x_argmax.unsqueeze(1)把形状从[3]变成[3,1]
  • 最终得到的是[3,1]的张量,用squeeze()去掉多余的维度就得到了我们需要的[3]形状结果

内容的提问来源于stack exchange,提问作者JVGD

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.08 11:32:52