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

PyTorch中根据给定索引张量提取对应位置张量值的方法

PyTorch二维张量按行匹配索引提取元素

问题描述

给定形状为(b, n)的值张量val、形状为(b, m)的索引张量ind(满足n>m),需要提取val中每一行内、ind对应列位置的元素。直接使用val[ind]做索引时,仅会扩展张量维度,无法得到行和索引一一对应的结果,输出形状不符合预期。

复现代码

import torch
val = torch.tensor([[1,2,3],
                    [4,5,6],
                    [7,8,9],
                    [10,11,12],
                    [13,14,15]])   
ind = torch.tensor([[1,2],
                    [0,2],
                    [0,1],
                    [1,2],
                    [0,1]])
val[ind] # 输出形状为(5,2,3),预期形状为(5,2)

预期输出

torch.tensor([[2,3],
              [4,6],
              [7,8],
              [11,12],
              [13,14]])

错误原因

直接传入ind做单参数索引时,PyTorch会将ind的所有值作为第0维(行维度)的索引,反复从val中抽取整行数据,最终输出形状为ind.shape + val.shape[1:],不会自动按行匹配索引位置。

正确实现方案

方案1:使用torch.gather(语义最清晰,推荐)

torch.gather是PyTorch内置的按索引取值API,指定沿列维度(dim=1)取对应索引位置的值即可:

result = torch.gather(val, dim=1, index=ind)

运行后得到的result形状为(5,2),和预期结果完全一致。

方案2:手动构造行索引做高级索引

先生成和ind形状一致的行坐标矩阵,再同时传入行、列两个索引做高级索引:

batch_size = val.shape[0]
# 构造行索引:每一行的元素都是当前行的行号,形状和ind一致
row_index = torch.arange(batch_size).unsqueeze(1).expand_as(ind)
result = val[row_index, ind]

两种方案输出结果完全相同,可根据代码场景选择使用。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 05:21:23