PyTorch中计算两个不等长一维整数张量的Jaccard相似度
PyTorch不等长一维整数张量Jaccard相似度实现
你现有的js_pt函数仅适用于等长二值多热编码向量的Jaccard距离计算,和你给出的基于集合操作的纯Pythonjs函数适用场景完全不同:前者要求输入长度等于总类别数、元素为0/1标记类别是否存在,后者输入是直接存储集合元素的整数序列,天然支持长度不等,因此不能直接复用原有等长版本的乘法求和逻辑。
下面是对齐纯Python集合逻辑、支持任意长度一维整数张量的PyTorch实现,全程使用PyTorch原生算子,兼容CPU/GPU设备,不需要转换为Python原生列表:
import torch def js_pt_unequal(t1: torch.Tensor, t2: torch.Tensor, empty_default: float = 0.0) -> torch.Tensor: # 对输入张量去重,对齐Python set的去重逻辑 s1 = torch.unique(t1) s2 = torch.unique(t2) # 广播对比两个集合的所有元素组合,统计交集元素个数 eq_mask = s1.unsqueeze(1) == s2.unsqueeze(0) intersection_cnt = eq_mask.any(dim=1).sum().float() # 按Jaccard公式计算并集大小 s1_len = s1.numel() s2_len = s2.numel() union_cnt = s1_len + s2_len - intersection_cnt # 处理两个集合全空的除零边界 if union_cnt == 0: return torch.tensor(empty_default, device=t1.device, dtype=torch.float32) # 返回Jaccard相似度,如果需要Jaccard距离就返回 1 - intersection_cnt / union_cnt return intersection_cnt / union_cnt
实现说明
- 逻辑完全对齐你提供的纯Python
js函数,对两个输入张量的长度没有相等要求,只要是一维整数张量即可 - 算子全部为PyTorch原生实现,张量不需要迁移到CPU转列表,支持GPU加速计算,自动匹配输入张量的设备
- 参数
empty_default用于处理两个输入全为空张量的边界场景,默认返回0.0,可根据业务需求调整为1.0(空集合视为完全相似) - 如果你需要的是Jaccard距离(即相似度取反,和你原有
js_pt返回值逻辑一致),只需要把最后返回语句替换为return 1 - intersection_cnt / union_cnt即可
测试示例
# 和纯Python实现结果对齐 a = torch.tensor([1,2,3,4,5]) b = torch.tensor([3,4,5,6,7,8]) print(js_pt_unequal(a, b)) # 输出0.375,和纯Python计算结果3/8完全一致
内容的提问来源于stack exchange,提问作者helloworld
相关产品推荐
相关产品推荐

