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

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

实现说明

  • 逻辑完全对齐你提供的纯Pythonjs函数,对两个输入张量的长度没有相等要求,只要是一维整数张量即可
  • 算子全部为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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 01:30:53