PyTorch反向索引:获取匹配reverse_indices的a元素对应索引
问题描述
现有两个PyTorch张量:
a: Tensor, shape (K,) reverse_indices: Tensor, shape (L,)
其中reverse_indices内所有值唯一且已排序,但a中的值不一定都存在于reverse_indices中。
需要生成一个LongTensor类型的张量indices(形状为(M,)),满足:
reverse_indices[indices] = a[torch.isin(a, reverse_indices)]
高效解决方案
由于reverse_indices是已排序且元素唯一的特性,我们可以利用二分查找来高效定位元素位置,避免低效的遍历匹配。具体步骤如下:
- 先筛选出
a中存在于reverse_indices的元素; - 使用
torch.searchsorted(基于二分查找)快速找到每个筛选后元素在reverse_indices中的索引位置,这就是目标indices。
代码示例
import torch # 示例输入 a = torch.tensor([3, 1, 4, 1, 5, 9, 2, 6]) reverse_indices = torch.tensor([1, 2, 3, 4, 5, 6]) # 筛选a中存在于reverse_indices的元素 a_filtered = a[torch.isin(a, reverse_indices)] # 二分查找获取对应索引 indices = torch.searchsorted(reverse_indices, a_filtered) # 验证正确性 assert torch.all(reverse_indices[indices] == a_filtered) print("生成的indices:", indices) # 输出:生成的indices: tensor([2, 0, 3, 0, 4, 5, 1, 5])
效率说明
torch.searchsorted的时间复杂度为O(M log L),其中M是a_filtered的长度,L是reverse_indices的长度。相比逐个元素遍历匹配的O(M*L)复杂度,在L较大的场景下性能提升非常明显。
内容的提问来源于stack exchange,提问作者tbrugere
相关产品推荐
相关产品推荐

