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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 06:11:03