PyTorch中高效查找两个大张量匹配行索引的方法
低显存实现PyTorch张量行匹配的索引查找
问题需求
给定两个2D CUDA张量A和B,需实现find_indices(A, B)函数,返回A中同时出现在B里的行的索引构成的1D张量,要求:
- 避免
for循环,支持并行化 - 内存占用≤1GB
已知前提条件:
- A和B的所有行均唯一;
- A和B的行均已排序,相同行在两者中顺序一致;
- 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
相关产品推荐
相关产品推荐

