如何统计Tensor元素在另一Tensor中的存在数量并优化效率?
问题描述
我有两个1D PyTorch张量:
A = torch.tensor([1, 2, 3, 4, 5, 6, 7, 8, 9, 10]) B = torch.tensor([2, 5, 6, 8, 12, 15, 16])
这些张量规模极大,长度不同,元素既非连续序列也未排序。我需要统计:
- B中存在于A的元素数量
- B中不存在于A的元素数量
预期输出:
Exists: 4 Do not exist: 3
尝试过的错误方法
我曾用以下代码:
exists = torch.eq(A,B).sum().item() not_exist = torch.numel(B) - exists
触发错误:
RuntimeError: The size of tensor a (10) must match the size of tensor b (7) at non-singleton dimension 0
可行但存疑的方法
下面的方法能运行,但需要先创建布尔张量再统计True的数量,想知道在处理超大张量时是否高效:
exists = np.isin(A,B).sum() not_exist = torch.numel(B) - exists
疑问
是否有更优或更高效的实现方案?
解决方案与分析
错误原因分析
torch.eq(A,B)要求两个张量形状完全匹配,只能逐元素比较对应位置的值,而你的A和B长度不同,所以直接报错。
现有numpy方法的效率问题
用np.isin(A,B)存在两个明显短板:
- 会强制把PyTorch张量转换成numpy数组,若张量原本在GPU上,会产生额外的CPU/GPU数据传输开销,超大张量场景下这部分成本很高。
- 即使在CPU上运行,跨框架转换的额外操作也会拖慢整体速度,不如原生PyTorch操作高效。
更高效的PyTorch原生方案
推荐使用PyTorch原生集合操作,避免跨框架转换,同时支持GPU加速,适合超大张量场景:
方法1:利用torch.isin(最优选择)
PyTorch 1.10及以上版本支持torch.isin,可直接在张量上执行元素存在性检查:
# 若A存在重复元素,先去重能减少后续匹配计算量 A_unique = torch.unique(A) # 生成B中元素是否在A中的布尔掩码 mask = torch.isin(B, A_unique) exists = mask.sum().item() not_exist = len(B) - exists print(f"Exists: {exists}") print(f"Do not exist: {not_exist}")
全程基于PyTorch张量操作,无需数据转换,GPU环境下能充分利用硬件加速,是超大张量场景的首选方案。若A本身无重复元素,可跳过torch.unique步骤。
方法2:哈希表手动实现(内存受限场景)
如果张量规模大到内存紧张,可先用哈希表存储A的元素,再遍历B统计:
# 将A转换为Python集合(若A在GPU上需先转CPU,需权衡内存占用) A_set = set(A.cpu().numpy()) exists = 0 for elem in B.cpu().numpy(): if elem in A_set: exists += 1 not_exist = len(B) - exists print(f"Exists: {exists}") print(f"Do not exist: {not_exist}")
这种方法内存占用更低,但如果张量原本在GPU上,数据传输的成本需要纳入考量,适合内存不足但CPU资源充足的场景。
效率对比总结
- 优先选择PyTorch原生
torch.isin,无跨框架转换开销,支持GPU加速,是超大张量场景下的最优解。 np.isin因涉及框架转换,效率低于原生PyTorch方法,不推荐使用。- 哈希表方法适合内存受限场景,但需注意GPU到CPU的数据传输成本。
内容的提问来源于stack exchange,提问作者bart
相关产品推荐
相关产品推荐

