PyTorch大张量场景下,如何矢量化获取张量B元素在A中的索引位置
PyTorch完全矢量化实现:为B中每个元素匹配A中的所有索引位置
问题背景
给定张量A(示例:torch.tensor([1,2,3,3,2,1,4,5,9]))和张量B(示例:torch.tensor([1, 2, 3, 9])),需要通过完全矢量化方法为B中每个元素找到其在A中的所有索引位置,输出需关联对应B元素(例如格式[[0,5], [1,4], [2,3], [-1,8]],一维/变长列表格式均可)。
原有方案存在的问题:
- 基于广播的矢量化函数在张量规模过大时会触发内存溢出,无法正常运行;
(A[..., None] == B).any(-1).nonzero()能获取匹配索引,但无法直接关联到B中对应的元素。
解决方案1:基于排序与二分查找的高效矢量化实现
该方法通过排序+二分查找避免大规模广播,时间复杂度为O(n log n),适合处理超大张量:
import torch from torch.nn.utils.rnn import pad_sequence def find_matching_indices(A, B): # 对A及其索引进行排序,为二分查找做准备 sorted_vals, sorted_indices = torch.sort(A) # 用二分查找定位B中每个元素在排序后A中的左右边界 left_bound = torch.searchsorted(sorted_vals, B, right=False) right_bound = torch.searchsorted(sorted_vals, B, right=True) # 为每个B元素提取对应的A索引,无匹配则返回[-1] result_list = [] for l, r in zip(left_bound, right_bound): if l == r: result_list.append(torch.tensor([-1], device=A.device)) else: result_list.append(sorted_indices[l:r]) # 可选:将变长列表转为固定长度张量,用-1填充空缺 return pad_sequence(result_list, batch_first=True, padding_value=-1)
测试示例
A = torch.tensor([1,2,3,3,2,1,4,5,9]) B = torch.tensor([1,2,3,9]) output = find_matching_indices(A, B) print(output)
输出:
tensor([[0, 5], [1, 4], [2, 3], [8, -1]])
解决方案2:基于唯一值映射的矢量化实现
该方法通过提取A的唯一值建立映射关系,适合A中重复值较多的场景:
import torch from torch.nn.utils.rnn import pad_sequence def find_matching_indices_v2(A, B): # 获取A的唯一值及对应逆映射 unique_vals, inverse_idx = torch.unique(A, return_inverse=True) # 预先生成唯一值到A索引的映射字典 val_to_indices = {} for idx, val in enumerate(unique_vals): val_to_indices[val.item()] = torch.where(inverse_idx == idx)[0] # 为B中每个元素匹配对应索引 result_list = [] for val in B: indices = val_to_indices.get(val.item(), torch.tensor([-1], device=A.device)) result_list.append(indices) # 可选:转为固定长度张量 return pad_sequence(result_list, batch_first=True, padding_value=-1)
方案说明
两种方案均为完全矢量化核心逻辑(仅外层对B的循环为轻量遍历,内部操作均为PyTorch矢量化运算),避免了大规模广播导致的内存爆炸问题,同时能准确关联B元素与对应A索引。若不需要固定长度输出,可直接返回result_list(变长张量列表)。
内容的提问来源于stack exchange,提问作者Andrew
相关产品推荐
相关产品推荐

