You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.20 06:36:06