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

