PyTorch批量计算向量Jaccard相似度的向量化实现方法
问题描述
现有一批形状为(bs, m, n)的张量,共bs个批次,每个批次包含m个长度为n的向量。需要针对每个批次,计算组内第一个向量和剩余m-1个向量的Jaccard相似度。
示例输入
a = [ [[3, 8, 6, 8, 7], [9, 7, 4, 8, 1], [7, 8, 8, 5, 7], [3, 9, 9, 4, 4]], [[7, 3, 8, 1, 7], [3, 0, 3, 4, 2], [9, 1, 6, 1, 6], [2, 7, 0, 6, 6]] ]
计算目标为a[:,0,:]与a[:,1:,:]的成对Jaccard相似度:第一批次中[3,8,6,8,7]分别和同批次后3个向量计算3个相似度得分,第二批次中[7,3,8,1,7]分别和同批次后3个向量计算3个相似度得分。
现有实现的问题
当前编写的单对向量Jaccard计算函数如下:
def js(la1, la2): combined = torch.cat((la1, la2)) union, counts = combined.unique(return_counts=True) intersection = union[counts > 1] torch.numel(intersection) / torch.numel(union)
该方法虽然能处理不等长张量,但每对向量组合的唯一值数量不一致,PyTorch不支持不规则张量,无法直接做批量处理。目前只能通过两层循环实现,计算效率很低:
bs = 2 m = 4 n = 5 a = torch.randint(0, 10, (bs, m, n)) print(f"Array is: \n{a}") for bs_idx in range(bs): first = a[bs_idx,0,:] for row in range(1, m): second = a[bs_idx,row,:] idx = js(first, second) print(f'comparing{first} and {second}: {idx}')
需要实现该逻辑的向量化版本,支持高效批量计算。
解决方案
如果向量元素取值范围是已知固定的(比如示例中取值为0-9的整数),可以直接用广播+计数的方式向量化实现,完全避免Python层循环,速度比循环版本快1~2个数量级。
实现思路
Jaccard相似度的核心计算公式为J = |A∩B| / |A∪B|,即两个集合交集大小除以并集大小:
- 先把每个向量转成multi-hot编码:对每个向量,标记哪些值出现过,不需要统计重复次数(集合去重后重复元素不影响交并集计算结果)
- 把每个批次第一个向量的multi-hot编码,和同批次其余向量的multi-hot编码做按位与,沿值维度求和得到交集大小
- 对两组编码做按位或,沿值维度求和得到并集大小
- 直接逐元素相除得到所有成对Jaccard相似度
向量化代码实现
import torch def batch_jaccard(a, value_range): """ 批量计算每个批次第一个向量和其余向量的Jaccard相似度 参数: a: 形状为(bs, m, n)的输入张量 value_range: 向量元素的取值总数,比如元素取0-9则传10 返回: 形状为(bs, m-1)的相似度矩阵 """ bs, m, n = a.shape # 生成multi-hot编码,形状(bs, m, value_range),对应位置为1表示该值在向量中出现过 one_hot = torch.zeros(bs, m, value_range, dtype=torch.bool, device=a.device) one_hot.scatter_(2, a.unsqueeze(-1), 1) # 取出每个批次第一个向量的编码,形状(bs, 1, value_range) first = one_hot[:, 0:1, :] # 取出每个批次其余向量的编码,形状(bs, m-1, value_range) rest = one_hot[:, 1:, :] # 计算交集大小:按位与后沿值维度求和 intersection = (first & rest).sum(dim=-1) # 计算并集大小:按位或后沿值维度求和 union = (first | rest).sum(dim=-1) # 加clamp避免除零,两个空向量比较时默认返回0 return intersection.float() / union.clamp(min=1).float() # 测试示例 if __name__ == "__main__": a = torch.tensor([ [[3, 8, 6, 8, 7], [9, 7, 4, 8, 1], [7, 8, 8, 5, 7], [3, 9, 9, 4, 4]], [[7, 3, 8, 1, 7], [3, 0, 3, 4, 2], [9, 1, 6, 1, 6], [2, 7, 0, 6, 6]] ]) # 示例中元素取值范围0-9,所以value_range传10 sim = batch_jaccard(a, value_range=10) print("批量计算得到的相似度:\n", sim)
结果验证
以示例输入为例,第一批次第一个向量[3,8,6,8,7]去重后集合为{3,6,7,8}:
- 和第二个向量
[9,7,4,8,1](去重集合{1,4,7,8,9})的交集大小为2、并集大小为7,相似度为2/7≈0.2857 - 和第三个向量
[7,8,8,5,7](去重集合{5,7,8})的交集大小为2、并集大小为5,相似度为2/5=0.4 - 和第四个向量
[3,9,9,4,4](去重集合{3,4,9})的交集大小为1、并集大小为6,相似度为1/6≈0.1667
计算结果和代码输出完全一致。
如果向量元素取值范围很大(比如超过1e5),multi-hot编码占用显存过高,可以先把所有元素映射到连续的最小取值区间再执行上述计算,整体效率依然远高于循环版本。
内容的提问来源于stack exchange,提问作者helloworld
相关产品推荐
相关产品推荐

