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

PyTorch中批量余弦相似度的高效实现(无for循环)

PyTorch高效计算批量余弦相似度(无for循环)

给定两个PyTorch张量:

  • a:形状为[batch_size, n, d],每个a[i,j]是d维向量
  • b:形状为[batch_size, m, d],每个b[i,k]是d维向量

需要输出张量c:形状为[batch_size, n, m],其中c[i,j,k]是a[i,j]与b[i,k]的余弦相似度,且不使用for循环实现高效计算。

实现思路

基于余弦相似度的核心公式:
$$\text{cos_sim}(u, v) = \frac{u \cdot v}{|u|_2 \times |v|_2}$$
利用PyTorch的批量矩阵运算和广播机制,直接完成批量内的相似度计算,完全规避循环。

代码实现

import torch

def batch_cosine_similarity(a, b, eps=1e-8):
    # 计算批量点积,得到[batch_size, n, m]的点积矩阵
    dot_product = torch.bmm(a, b.transpose(1, 2))
    # 计算a中每个向量的L2范数,保持维度为[batch_size, n, 1]
    norm_a = torch.norm(a, dim=2, keepdim=True)
    # 计算b中每个向量的L2范数,转置后维度为[batch_size, 1, m]
    norm_b = torch.norm(b, dim=2, keepdim=True).transpose(1, 2)
    # 计算分母,添加eps避免除以0的情况
    denominator = norm_a * norm_b + eps
    # 逐元素相除得到余弦相似度矩阵
    cos_sim = dot_product / denominator
    return cos_sim

关键步骤解释

  1. 批量点积计算:使用torch.bmm(批量矩阵乘法),将b的后两维转置为[batch_size, d, m],与a([batch_size, n, d])相乘,直接得到每个batch内n×m的点积结果。
  2. 范数与广播:用torch.norm指定dim=2计算每个d维向量的L2范数,keepdim=True保留维度以便后续广播;将b的范数张量转置后,和a的范数张量通过广播机制相乘,得到与点积矩阵同形状的分母矩阵。
  3. 除零防护:添加极小值eps,避免因向量模长为0导致的除以0错误。

测试示例

# 构造测试张量
batch_size = 2
n = 3
m = 4
d = 5

a = torch.randn(batch_size, n, d)
b = torch.randn(batch_size, m, d)

# 计算余弦相似度
c = batch_cosine_similarity(a, b)
print(c.shape)  # 输出: torch.Size([2, 3, 4])

内容的提问来源于stack exchange,提问作者helperFunction

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 00:05:26