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

PyTorch中高效查找两个大张量匹配行索引的方法

低显存实现PyTorch张量行匹配的索引查找

问题需求

给定两个2D CUDA张量A和B,需实现find_indices(A, B)函数,返回A中同时出现在B里的行的索引构成的1D张量,要求:

  • 避免for循环,支持并行化
  • 内存占用≤1GB

已知前提条件:

  1. A和B的所有行均唯一;
  2. A和B的行均已排序,相同行在两者中顺序一致;
  3. A和B的行数约为20万。

示例代码:

import torch

A = torch.tensor([[1, 2, 3], [2, 3, 4], [3, 4, 5]]).cuda()
B = torch.tensor([[1, 2, 3], [2, 3, 6], [2, 5, 6], [3, 4, 5]]).cuda()

indices1 = find_indices(A, B)  # 预期输出:tensor([0, 2])
indices2 = find_indices(B, A)  # 预期输出:tensor([0, 3])

assert A[indices1].equal(B[indices2])

之前尝试的方法因显存占用超100GB触发CUDA内存不足:

values, indices = torch.topk(((A.t() == B.unsqueeze(-1)).all(dim=1)).int(), 1, 1)
indices = indices[values!=0]

解决方案

利用行已排序、唯一的前提条件,提供两种低显存实现方式:

方法1:基于唯一键的二分查找法

将每行转换为唯一标量键,再通过二分查找快速匹配,内存占用极低:

def find_indices(A, B):
    # 计算基数,确保大于所有元素值,避免键冲突
    max_val = torch.max(torch.cat([A, B]))
    base = max_val + 1
    dim = A.shape[1]
    
    # 生成权重:base^0, base^1, ..., base^(dim-1)
    weights = base ** torch.arange(dim, device=A.device, dtype=torch.int64)
    
    # 将每行转换为唯一标量键
    A_keys = torch.sum(A * weights, dim=1)
    B_keys = torch.sum(B * weights, dim=1)
    
    # 查找A中每个键在B中的位置
    pos_in_B = torch.searchsorted(B_keys, A_keys)
    # 过滤出确实存在于B中的行的索引
    mask = (pos_in_B < len(B_keys)) & (B_keys[pos_in_B] == A_keys)
    
    return torch.where(mask)[0]

优势:

  • 内存占用极小:仅需存储两个1D张量(20万元素/个),总内存约3.2MB(int64类型)
  • 时间复杂度O(n log m),20万规模下计算极快
  • 无哈希冲突,结果绝对准确

方法2:基于拼接排序的重复行检测法

通过拼接张量并排序,快速定位重复行,同样低显存:

def find_indices(A, B):
    # 为A添加原索引,为B添加标记(-1表示来自B)
    A_with_idx = torch.cat([A, torch.arange(len(A), device=A.device).unsqueeze(1)], dim=1)
    B_with_idx = torch.cat([B, torch.full((len(B), 1), -1, device=B.device)], dim=1)
    
    # 拼接两个张量
    combined = torch.cat([A_with_idx, B_with_idx], dim=0)
    # 按行的字典序排序(利用原行已排序的特性)
    sorted_combined = combined[torch.lexsort(combined[:, :-1].t())]
    
    # 检测相邻行是否为重复行,且分别来自A和B
    mask = torch.all(sorted_combined[:-1, :-1] == sorted_combined[1:, :-1], dim=1)
    mask &= (sorted_combined[:-1, -1] != -1) & (sorted_combined[1:, -1] == -1)
    
    # 提取A中匹配行的原索引
    return sorted_combined[:-1, -1][mask].long()

优势:

  • 内存占用低:拼接后张量约12.8MB(按每行4个int64元素计算)
  • 无需计算基数,避免大数溢出风险(当元素值极大时更安全)
  • 结果准确,符合需求

验证

运行示例代码测试:

A = torch.tensor([[1, 2, 3], [2, 3, 4], [3, 4, 5]]).cuda()
B = torch.tensor([[1, 2, 3], [2, 3, 6], [2, 5, 6], [3, 4, 5]]).cuda()

indices1 = find_indices(A, B)
indices2 = find_indices(B, A)

assert indices1.equal(torch.tensor([0, 2], device='cuda'))
assert indices2.equal(torch.tensor([0, 3], device='cuda'))
assert A[indices1].equal(B[indices2])

内容的提问来源于stack exchange,提问作者Eureka D

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 03:58:13