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
相关产品推荐
相关产品推荐

