如何基于索引张量从原张量提取元素并保留原维度顺序
问题描述
我有一个形状为(1,3,1)的张量:
topk_indices = tensor([[[6], [1], [0]]], device='cuda:0')
以及一个形状为(1,7,16)的张量:
in_tensor = tensor([[[0.8359, 0.4812, 0.0297, 0.5219, 0.1595, 0.9066, 0.1965, 0.4639, 0.3890, 0.5890, 0.9705, 0.5475, 0.7896, 0.8881, 0.9037, 0.3273], [0.3882, 0.7410, 0.3636, 0.7341, 0.3908, 0.1609, 0.7035, 0.5767, 0.7229, 0.9967, 0.8414, 0.9740, 0.5268, 0.0699, 0.1492, 0.1894], [0.0594, 0.2494, 0.0397, 0.0387, 0.2012, 0.0071, 0.1931, 0.6907, 0.9170, 0.3513, 0.3546, 0.7670, 0.2533, 0.2636, 0.8081, 0.0643], [0.5611, 0.9417, 0.5857, 0.6360, 0.2088, 0.4931, 0.5275, 0.6227, 0.6943, 0.9345, 0.1184, 0.5150, 0.2502, 0.1045, 0.4600, 0.0599], [0.8489, 0.5579, 0.2305, 0.7613, 0.0268, 0.3066, 0.4026, 0.0751, 0.1821, 0.4184, 0.8794, 0.9828, 0.8181, 0.2014, 0.1729, 0.9363], [0.6769, 0.5133, 0.5677, 0.0982, 0.3331, 0.9813, 0.3767, 0.4749, 0.0848, 0.2203, 0.4898, 0.1894, 0.4380, 0.7035, 0.0109, 0.6485], [0.1694, 0.2560, 0.6920, 0.8976, 0.3633, 0.2947, 0.0479, 0.2422, 0.0622, 0.3856, 0.6020, 0.0316, 0.9366, 0.8137, 0.0105, 0.2612]]], device='cuda:0')
我想要生成一个形状为(1,3,16)的新张量,只保留topk_indices里指定的索引对应的元素,但要按照in_tensor在dim=1维度上的原始顺序排列。期望结果如下:
in_tensor = tensor([[[0.8359, 0.4812, 0.0297, 0.5219, 0.1595, 0.9066, 0.1965, 0.4639, 0.3890, 0.5890, 0.9705, 0.5475, 0.7896, 0.8881, 0.9037, 0.3273], [0.3882, 0.7410, 0.3636, 0.7341, 0.3908, 0.1609, 0.7035, 0.5767, 0.7229, 0.9967, 0.8414, 0.9740, 0.5268, 0.0699, 0.1492, 0.1894], [0.1694, 0.2560, 0.6920, 0.8976, 0.3633, 0.2947, 0.0479, 0.2422, 0.0622, 0.3856, 0.6020, 0.0316, 0.9366, 0.8137, 0.0105, 0.2612]]], device='cuda:0')
也就是保留索引6、1、0对应的元素,但按原张量里的顺序(0、1、6)排列。我试过用torch.gather:
selected_tensor = torch.gather(in_tensor, 1, topk_indices.repeat(1, 1, in_tensor.shape[-1]).unsqueeze(3).squeeze(3))
但得到的结果是按topk_indices的顺序排列的,不符合需求,求正确实现方法。
解决方案
核心思路是先提取topk_indices中的索引值,对这些索引排序后,直接用排序后的索引去选取in_tensor对应维度的元素,就能保证结果是原张量中的顺序。
具体代码实现:
# 提取索引并排序,得到原张量中的顺序索引 sorted_indices = topk_indices.squeeze().sort()[0] # 选取对应维度的元素 selected_tensor = in_tensor[:, sorted_indices, :]
说明
topk_indices.squeeze()会把形状从(1,3,1)压缩为(3,)的一维张量,方便后续排序操作;.sort()[0]会返回排序后的索引序列,这里得到的是tensor([0, 1, 6], device='cuda:0');- 用排序后的索引直接索引
in_tensor的第二维度(dim=1),就能得到形状为(1,3,16)的目标张量,且元素顺序和原张量中0、1、6位置的顺序完全一致。
这种方法比torch.gather更简洁直接,也能完美满足需求。
内容的提问来源于stack exchange,提问作者PatrickHellman
相关产品推荐
相关产品推荐

