PyTorch中如何根据标签张量提取得分张量对应索引位置的值
PyTorch对应实现方法
PyTorch提供了直接实现该需求的内置功能,以下是两种常用的实现方案,均支持自动梯度计算:
方案1:使用torch.gather()内置函数
这是专门用于按指定索引从张量中提取值的API,适配任意维度的张量提取场景,你的场景下用法如下:
# y的形状为[5],先扩展为[5,1]和scores的二维结构匹配 selected = scores.gather(dim=1, index=y.unsqueeze(dim=1)).squeeze(dim=1)
参数说明:
dim=1:指定在scores的第1维(列维度)上提取元素index=y.unsqueeze(dim=1):索引张量需要和输入张量的维度数一致,所以给y额外加一个维度- 最后用
squeeze去掉多余的维度,得到形状为[5]的提取结果
方案2:使用高级索引写法
对于二维张量的行对齐提取场景,可以用更简洁的高级索引语法实现,不需要调用额外API:
# 第一个索引取0-4的行序号,第二个索引取y中对应的列号,逐行匹配取值 selected = scores[torch.arange(scores.shape[0]), y]
该写法直接得到形状为[5]的结果,逻辑更直观。
内容的提问来源于stack exchange,提问作者Arneo
相关产品推荐
相关产品推荐

