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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 03:45:03