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

