如何在PyTorch中实现Numpy带return_index=True的unique功能?
PyTorch实现类似numpy.unique(return_index=True)及扩展场景的方案
一、获取张量中唯一元素的首次出现索引
PyTorch原生torch.unique没有return_index参数,我们可以通过以下两种方式实现类似功能:
通用兼容方案(支持所有可排序类型)
通过排序+掩码标记的方式实现,兼容整数、浮点数等任意可排序张量:
def unique_with_indices(tensor): # 排序张量并记录原索引 sorted_tensor, sorted_indices = torch.sort(tensor) # 生成掩码:标记唯一元素的位置(第一个元素默认保留,后续元素与前一个不同则保留) unique_mask = torch.cat([torch.tensor([True], device=tensor.device), sorted_tensor[1:] != sorted_tensor[:-1]]) # 提取唯一元素对应的原索引和元素值 unique_indices = sorted_indices[unique_mask] unique_elements = sorted_tensor[unique_mask] # 按原张量中首次出现的顺序重新排列结果 sorted_order = torch.argsort(unique_indices) return unique_elements[sorted_order], unique_indices[sorted_order]
优化方案(仅适用于非负整数张量)
利用哈希标记避免排序操作,效率更高,但有元素范围限制:
def unique_with_indices_int(tensor): # 仅支持非负整数张量 max_val = tensor.max().item() # 初始化标记数组,-1表示未见过该元素 seen = torch.full((max_val + 1,), -1, dtype=torch.long, device=tensor.device) indices = torch.arange(tensor.size(0), device=tensor.device) # 标记首次出现的元素索引 mask = seen[tensor] == -1 seen[tensor[mask]] = indices[mask] # 提取唯一元素和对应索引 unique_elements = torch.nonzero(seen != -1).squeeze() unique_indices = seen[seen != -1] # 按首次出现顺序排序 sorted_order = torch.argsort(unique_indices) return unique_elements[sorted_order], unique_indices[sorted_order]
二、获取v2中不在v1内的元素的首次出现索引
通用场景实现
先筛选v2中不属于v1的元素,再提取这些元素的首次出现原索引:
def get_unique_indices_not_in_v1(v1, v2): # 获取v1的唯一元素集合,减少重复判断 v1_unique = torch.unique(v1) # 标记v2中不在v1里的元素 mask = ~torch.isin(v2, v1_unique) filtered_v2 = v2[mask] filtered_indices = torch.arange(v2.size(0), device=v2.device)[mask] if filtered_v2.numel() == 0: return torch.tensor([], dtype=torch.long, device=v2.device) # 提取筛选后元素的首次出现索引(对应原v2的位置) _, unique_filtered_idx = unique_with_indices(filtered_v2) original_indices = filtered_indices[unique_filtered_idx] # 按原v2中的出现顺序排序结果 original_indices, _ = torch.sort(original_indices) return original_indices
已知v1已排序时的优化方案
当v1是升序排列的张量时,用torch.searchsorted替代torch.isin可大幅提升判断效率:
def get_unique_indices_not_in_sorted_v1(sorted_v1, v2): # 前提:sorted_v1必须为升序排列 # 查找v2元素在sorted_v1中的插入位置 pos = torch.searchsorted(sorted_v1, v2) # 判断元素是否不在v1中:插入位置超出范围,或对应位置元素不匹配 mask = (pos >= sorted_v1.size(0)) | (sorted_v1[pos] != v2) filtered_v2 = v2[mask] filtered_indices = torch.arange(v2.size(0), device=v2.device)[mask] if filtered_v2.numel() == 0: return torch.tensor([], dtype=torch.long, device=v2.device) _, unique_filtered_idx = unique_with_indices(filtered_v2) original_indices = filtered_indices[unique_filtered_idx] original_indices, _ = torch.sort(original_indices) return original_indices
补充说明
- 排序实现的方法兼容性强,但会引入排序的时间开销;
- 整数哈希方法效率更高,但仅支持非负整数,且元素最大值不能过大(否则标记数组会占用过多内存);
- 当v1已排序时,
searchsorted的时间复杂度为O(n log m)(n为v2长度,m为v1长度),比isin的O(n + m log m)更高效,适合大张量场景。
内容的提问来源于stack exchange,提问作者Alexander Chervov
相关产品推荐
相关产品推荐

