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

如何在PyTorch中根据张量最大值索引提取另一张量对应值

解决PyTorch中根据张量最大值索引提取对应张量值的问题

直接使用b[indices]不符合预期,是因为indices的维度为(1,3),而b的维度是(2,3),直接索引会将indices中的每个元素视为行索引,取出整行后堆叠,得到(1,3,3)的张量,并非目标结果。

下面提供两种可行的解决方法:

方法一:使用torch.gather(推荐)

torch.gather是PyTorch专门用于按指定维度和索引提取元素的函数,适配这类场景:

import torch

a = torch.tensor([[1,2,4],[2,1,3]])
b = torch.tensor([[10,24,2],[23,4,5]])

max_values, indices = torch.max(a, dim=0, keepdim=True)
# 在dim=0维度上按indices提取b的对应元素
result = torch.gather(b, dim=0, index=indices)
print(result)
# 输出:tensor([[23, 24, 2]])

由于torch.max设置了keepdim=True,indices的形状与b在dim=0外的维度完全匹配,可直接传入gather完成提取。

方法二:使用高级索引手动匹配

通过生成对应列索引,配合行索引indices进行精准定位:

import torch

a = torch.tensor([[1,2,4],[2,1,3]])
b = torch.tensor([[10,24,2],[23,4,5]])

max_values, indices = torch.max(a, dim=0, keepdim=True)
# 生成与indices形状一致的列索引
col_indices = torch.arange(b.size(1)).unsqueeze(0)
# 行索引与列索引一一对应提取元素
result = b[indices, col_indices]
print(result)
# 输出:tensor([[23, 24, 2]])

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 18:35:11