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

如何使用torch.max返回的索引从另一张量中提取对应位置的元素

PyTorch使用max返回索引取其他张量对应元素的实现方法

你可以使用torch.gather接口实现取值,该接口专门用于按照指定维度的索引从张量中提取对应元素,实现代码如下:

import torch
a = torch.rand(2,3,4)
b = torch.rand(2,3,4)
# 取dim=1维度上的最大值索引,shape为[2,4]
indices = torch.max(a, 1)[1]
# 给索引增加dim=1维度,和b的3维结构匹配,取值后去掉多余维度
b_max = torch.gather(b, dim=1, index=indices.unsqueeze(1)).squeeze(1)
  • 维度匹配说明:torch.gather要求传入的index张量维度数必须和待取值张量的维度数完全一致,因此需要先对shape为[2,4]的indices调用unsqueeze(1)扩展为[2,1,4],和b的[2,3,4]维度数匹配,取值完成后调用squeeze(1)去掉多余的维度,最终得到的b_max形状为[2,4],和indices原始形状一致。

如果你更习惯用高级索引写法,也可以用如下方式实现,和torch.gather效果完全一致:

batch_size, _, feat_dim = a.shape
# 构造对应维度的索引,和indices形状对齐
batch_idx = torch.arange(batch_size)[:, None].expand(-1, feat_dim)
feat_idx = torch.arange(feat_dim)[None, :].expand(batch_size, -1)
b_max = b[batch_idx, indices, feat_idx]

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 14:39:03