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
关键步骤解释
- 批量点积计算:使用
torch.bmm(批量矩阵乘法),将b的后两维转置为[batch_size, d, m],与a([batch_size, n, d])相乘,直接得到每个batch内n×m的点积结果。 - 范数与广播:用
torch.norm指定dim=2计算每个d维向量的L2范数,keepdim=True保留维度以便后续广播;将b的范数张量转置后,和a的范数张量通过广播机制相乘,得到与点积矩阵同形状的分母矩阵。 - 除零防护:添加极小值
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
相关产品推荐
相关产品推荐

