【Bug已解决】How to compute the cosine_similarity in pytorch for all rows in a matrix with respect to all rows in another matrix 解决方案问题描述在深度学习和自然语言处理中计算两个矩阵之间所有行的余弦相似度cosine similarity是一个常见需求。例如在句子相似度计算、推荐系统、对比学习contrastive learning等场景中需要计算一个矩阵中每一行与另一个矩阵中每一行的余弦相似度得到一个 N×M 的相似度矩阵。许多开发者不知道如何高效地使用 PyTorch 实现这个操作常见的问题包括如何用矩阵运算代替循环计算余弦相似度F.cosine_similarity的默认行为是逐行计算如何改为全交叉计算如何避免数值不稳定除零问题如何处理大规模矩阵的内存问题如何在 GPU 上高效计算torch.nn.functional.cosine_similarity的默认行为是对两个张量的对应行计算余弦相似度逐行配对而不是全交叉计算。要实现全交叉计算需要使用广播broadcasting技巧或矩阵乘法。错误复现场景一误用 F.cosine_similarityimport torch import torch.nn.functional as F # 两个矩阵 A torch.randn(5, 128) # 5 个向量每个 128 维 B torch.randn(3, 128) # 3 个向量每个 128 维 # 尝试计算余弦相似度 try: sim F.cosine_similarity(A, B) except RuntimeError as e: print(f错误: {e}) # The size of tensor a (5) must match the size of tensor b (3)场景二使用循环效率低# 使用循环计算 - 正确但效率低 A torch.randn(1000, 128) B torch.randn(500, 128) sim_matrix torch.zeros(1000, 500) for i in range(1000): for j in range(500): sim_matrix[i, j] F.cosine_similarity(A[i], B[j], dim0) # 正确但非常慢1000x500 50万次循环场景三除零错误# 当向量全为零时余弦相似度计算会除零 A torch.zeros(1, 128) B torch.randn(1, 128) # 手动计算 dot_product (A * B).sum(dim1) norm_a A.norm(dim1) norm_b B.norm(dim1) cos_sim dot_product / (norm_a * norm_b) print(cos_sim) # nan! 除零场景四广播维度错误A torch.randn(5, 128) B torch.randn(3, 128) # 尝试使用广播 - 维度不对 try: sim F.cosine_similarity(A.unsqueeze(1), B.unsqueeze(0), dim2) except RuntimeError as e: print(f错误: {e}) # 维度不匹配根因分析1. F.cosine_similarity 的默认行为F.cosine_similarity(x1, x2, dim1)计算的是逐行配对的余弦相似度# x1: (N, D), x2: (N, D) # 输出: (N,) - 第 i 个输出是 x1[i] 和 x2[i] 的余弦相似度 sim F.cosine_similarity(x1, x2, dim1)它要求x1和x2在非dim维度上形状相同不支持全交叉计算。2. 全交叉计算的需求全交叉计算需要得到 N×M 的矩阵其中sim[i, j]是A[i]和B[j]的余弦相似度。这需要使用广播或矩阵乘法来实现。3. 余弦相似度的数学定义余弦相似度公式cos_sim(a, b) (a · b) / (||a|| * ||b||)对于矩阵 A (N×D) 和 B (M×D)全交叉余弦相似度矩阵为sim[i, j] sum(A[i, :] * B[j, :]) / (norm(A[i, :]) * norm(B[j, :]))可以用矩阵乘法高效计算sim (A_normalized B_normalized.T)其中A_normalized和B_normalized是 L2 归一化后的矩阵。4. 数值稳定性问题当向量的范数为 0 时全零向量除法会产生 NaN。需要添加小的 epsilon 值来避免除零。解决方案方案一使用矩阵乘法推荐最高效import torch import torch.nn.functional as F def cosine_similarity_matrix(A: torch.Tensor, B: torch.Tensor, eps: float 1e-8) - torch.Tensor: 计算两个矩阵所有行之间的余弦相似度 Args: A: (N, D) 张量 B: (M, D) 张量 eps: 数值稳定性的小值 Returns: (N, M) 余弦相似度矩阵 # L2 归一化 A_norm A / (A.norm(dim1, keepdimTrue) eps) B_norm B / (B.norm(dim1, keepdimTrue) eps) # 矩阵乘法 sim A_norm B_norm.T # (N, M) return sim # 使用示例 A torch.randn(5, 128) B torch.randn(3, 128) sim cosine_similarity_matrix(A, B) print(f相似度矩阵形状: {sim.shape}) # (5, 3) print(f相似度矩阵:\n{sim})方案二使用 F.normalize 矩阵乘法def cosine_similarity_matrix_v2(A: torch.Tensor, B: torch.Tensor, eps: float 1e-8) - torch.Tensor: 使用 F.normalize 计算余弦相似度矩阵 # 使用 F.normalize 进行 L2 归一化 A_norm F.normalize(A, p2, dim1, epseps) B_norm F.normalize(B, p2, dim1, epseps) # 矩阵乘法 sim A_norm B_norm.T return sim # 使用 A torch.randn(5, 128) B torch.randn(3, 128) sim cosine_similarity_matrix_v2(A, B) print(f相似度矩阵形状: {sim.shape}) # (5, 3)方案三使用广播 F.cosine_similaritydef cosine_similarity_matrix_v3(A: torch.Tensor, B: torch.Tensor) - torch.Tensor: 使用广播和 F.cosine_similarity 计算 适用于较小的矩阵 # A: (N, D) - (N, 1, D) # B: (M, D) - (1, M, D) # 广播后: (N, M, D) A_expanded A.unsqueeze(1) # (N, 1, D) B_expanded B.unsqueeze(0) # (1, M, D) # F.cosine_similarity 在 dim2 上计算 sim F.cosine_similarity(A_expanded, B_expanded, dim2) # (N, M) return sim # 使用 A torch.randn(5, 128) B torch.randn(3, 128) sim cosine_similarity_matrix_v3(A, B) print(f相似度矩阵形状: {sim.shape}) # (5, 3)方案四手动实现带数值稳定性def cosine_similarity_matrix_v4(A: torch.Tensor, B: torch.Tensor, eps: float 1e-8) - torch.Tensor: 手动实现余弦相似度矩阵带数值稳定性处理 # 计算点积矩阵: (N, M) dot_product A B.T # 计算范数: (N,) 和 (M,) norm_A A.norm(dim1, keepdimTrue) # (N, 1) norm_B B.norm(dim1, keepdimTrue) # (M, 1) # 外积得到范数矩阵: (N, M) norm_matrix norm_A norm_B.T # 余弦相似度 sim dot_product / (norm_matrix eps) # 处理 NaN全零向量情况 sim torch.nan_to_num(sim, nan0.0) return sim # 使用 A torch.randn(5, 128) B torch.randn(3, 128) sim cosine_similarity_matrix_v4(A, B) print(f相似度矩阵形状: {sim.shape})方案五GPU 加速版本def cosine_similarity_gpu(A: torch.Tensor, B: torch.Tensor, device: str cuda, eps: float 1e-8) - torch.Tensor: GPU 加速的余弦相似度计算 支持大规模矩阵 A A.to(device) B B.to(device) A_norm F.normalize(A, p2, dim1, epseps) B_norm F.normalize(B, p2, dim1, epseps) sim A_norm B_norm.T return sim # 使用 if torch.cuda.is_available(): A torch.randn(10000, 768) B torch.randn(5000, 768) sim cosine_similarity_gpu(A, B) print(fGPU 相似度矩阵形状: {sim.shape}) # (10000, 5000)完整修复代码以下是一个完整的余弦相似度计算工具模块 PyTorch 余弦相似度矩阵计算工具 支持全交叉计算、批处理、GPU 加速 import torch import torch.nn.functional as F from typing import Optional, Tuple import time class CosineSimilarity: 余弦相似度计算工具类 staticmethod def compute(A: torch.Tensor, B: torch.Tensor, eps: float 1e-8) - torch.Tensor: 计算两个矩阵所有行之间的余弦相似度 Args: A: (N, D) 张量 B: (M, D) 张量 eps: 数值稳定性的小值 Returns: (N, M) 余弦相似度矩阵值域 [-1, 1] A_norm F.normalize(A, p2, dim1, epseps) B_norm F.normalize(B, p2, dim1, epseps) return A_norm B_norm.T staticmethod def compute_batched(A: torch.Tensor, B: torch.Tensor, batch_size: int 1024, eps: float 1e-8) - torch.Tensor: 分批计算余弦相似度节省内存 Args: A: (N, D) 张量 B: (M, D) 张量 batch_size: 每批处理的行数 eps: 数值稳定性的小值 Returns: (N, M) 余弦相似度矩阵 N A.size(0) M B.size(0) B_norm F.normalize(B, p2, dim1, epseps) result torch.zeros(N, M, deviceA.device, dtypeA.dtype) for start in range(0, N, batch_size): end min(start batch_size, N) A_batch A[start:end] A_batch_norm F.normalize(A_batch, p2, dim1, epseps) result[start:end] A_batch_norm B_norm.T return result staticmethod def top_k_similar(A: torch.Tensor, B: torch.Tensor, k: int 5, eps: float 1e-8) - Tuple[torch.Tensor, torch.Tensor]: 找到 A 中每行与 B 中最相似的 k 行 Args: A: (N, D) 张量 B: (M, D) 张量 k: 返回的 top-k 数量 eps: 数值稳定性的小值 Returns: (values, indices) - 相似度值和对应的 B 中索引 sim CosineSimilarity.compute(A, B, eps) values, indices sim.topk(k, dim1, largestTrue) return values, indices staticmethod def pairwise_distance(A: torch.Tensor, B: torch.Tensor, eps: float 1e-8) - torch.Tensor: 计算余弦距离矩阵 (1 - cosine_similarity) Returns: (N, M) 距离矩阵值域 [0, 2] sim CosineSimilarity.compute(A, B, eps) return 1.0 - sim staticmethod def self_similarity(A: torch.Tensor, eps: float 1e-8) - torch.Tensor: 计算矩阵自身的余弦相似度矩阵 Args: A: (N, D) 张量 Returns: (N, N) 余弦相似度矩阵 return CosineSimilarity.compute(A, A, eps) staticmethod def compute_with_temperature(A: torch.Tensor, B: torch.Tensor, temperature: float 0.07, eps: float 1e-8) - torch.Tensor: 带温度系数的余弦相似度用于对比学习 Args: A: (N, D) 张量 B: (M, D) 张量 temperature: 温度系数 eps: 数值稳定性的小值 Returns: (N, M) 缩放后的相似度矩阵 sim CosineSimilarity.compute(A, B, eps) return sim / temperature # # 完整示例 # def demo_basic(): 基本用法演示 print( * 60) print(基本余弦相似度计算) print( * 60) A torch.randn(5, 128) B torch.randn(3, 128) sim CosineSimilarity.compute(A, B) print(fA 形状: {A.shape}) print(fB 形状: {B.shape}) print(f相似度矩阵形状: {sim.shape}) # (5, 3) print(f相似度矩阵:\n{sim}) print(f值域: [{sim.min():.4f}, {sim.max():.4f}]) def demo_top_k(): Top-K 相似搜索 print(\n * 60) print(Top-K 相似搜索) print( * 60) # 模拟句子嵌入 A torch.randn(4, 768) # 4 个查询句子 B torch.randn(100, 768) # 100 个候选句子 values, indices CosineSimilarity.top_k_similar(A, B, k3) for i in range(A.size(0)): print(f\n查询 {i1} 的 Top-3 相似:) for j in range(3): print(f 候选 {indices[i, j].item()}: f相似度{values[i, j].item():.4f}) def demo_contrastive_learning(): 对比学习中的余弦相似度 print(\n * 60) print(对比学习中的余弦相似度) print( * 60) # SimCLR 风格的对比学习 batch_size 8 feature_dim 128 # 两个增强视图的嵌入 z_i F.normalize(torch.randn(batch_size, feature_dim), dim1) z_j F.normalize(torch.randn(batch_size, feature_dim), dim1) # 计算相似度矩阵带温度系数 temperature 0.5 sim_matrix CosineSimilarity.compute_with_temperature( z_i, z_j, temperaturetemperature ) print(fz_i 形状: {z_i.shape}) print(fz_j 形状: {z_j.shape}) print(f相似度矩阵形状: {sim_matrix.shape}) # (8, 8) print(f温度系数: {temperature}) # 对比损失 (InfoNCE) labels torch.arange(batch_size) loss F.cross_entropy(sim_matrix, labels) print(f对比损失: {loss.item():.4f}) def demo_large_scale(): 大规模矩阵计算 print(\n * 60) print(大规模矩阵计算) print( * 60) N, M, D 5000, 3000, 512 A torch.randn(N, D) B torch.randn(M, D) # 直接计算 print(f矩阵大小: A({N}x{D}), B({M}x{D})) print(f输出矩阵大小: {N}x{M} {N*M/1e6:.1f}M 元素) start time.time() sim CosineSimilarity.compute(A, B) direct_time time.time() - start print(f直接计算耗时: {direct_time:.4f}s) # 分批计算 start time.time() sim_batched CosineSimilarity.compute_batched(A, B, batch_size1000) batched_time time.time() - start print(f分批计算耗时: {batched_time:.4f}s) # 验证结果一致 print(f结果一致: {torch.allclose(sim, sim_batched, atol1e-5)}) def demo_edge_cases(): 边界情况处理 print(\n * 60) print(边界情况处理) print( * 60) # 全零向量 A torch.tensor([[0.0, 0.0, 0.0], [1.0, 2.0, 3.0]]) B torch.tensor([[0.0, 0.0, 0.0], [4.0, 5.0, 6.0]]) sim CosineSimilarity.compute(A, B) print(f含全零向量的相似度矩阵:\n{sim}) print(fNaN 数量: {torch.isnan(sim).sum().item()}) # 相同向量 A torch.tensor([[1.0, 2.0, 3.0]]) B torch.tensor([[1.0, 2.0, 3.0]]) sim CosineSimilarity.compute(A, B) print(f\n相同向量的相似度: {sim.item():.6f}) # 应该接近 1.0 # 相反向量 A torch.tensor([[1.0, 2.0, 3.0]]) B torch.tensor([[-1.0, -2.0, -3.0]]) sim CosineSimilarity.compute(A, B) print(f相反向量的相似度: {sim.item():.6f}) # 应该接近 -1.0 def demo_performance_comparison(): 性能对比 print(\n * 60) print(性能对比: 循环 vs 矩阵乘法) print( * 60) N, M, D 100, 50, 128 A torch.randn(N, D) B torch.randn(M, D) # 方法 1: 循环 start time.time() sim_loop torch.zeros(N, M) for i in range(N): for j in range(M): sim_loop[i, j] F.cosine_similarity( A[i].unsqueeze(0), B[j].unsqueeze(0) ) loop_time time.time() - start # 方法 2: 矩阵乘法 start time.time() sim_matrix CosineSimilarity.compute(A, B) matrix_time time.time() - start print(f循环方法: {loop_time:.4f}s) print(f矩阵乘法: {matrix_time:.4f}s) print(f加速比: {loop_time / matrix_time:.1f}x) print(f结果一致: {torch.allclose(sim_loop, sim_matrix, atol1e-5)}) if __name__ __main__: demo_basic() demo_top_k() demo_contrastive_learning() demo_large_scale() demo_edge_cases() demo_performance_comparison()常见陷阱与注意事项1. F.cosine_similarity 是逐行配对不是全交叉# 逐行配对默认 A torch.randn(5, 128) B torch.randn(5, 128) sim F.cosine_similarity(A, B, dim1) # (5,) - A[i] 和 B[i] 的相似度 # 全交叉计算 sim_matrix CosineSimilarity.compute(A, B) # (5, 5) - 所有配对2. 数值稳定性# 全零向量会导致除零 A torch.zeros(1, 128) B torch.randn(1, 128) # 错误: 直接除法 norm_a A.norm(dim1) # 0 sim (A B.T) / (norm_a * B.norm(dim1)) # nan! # 正确: 使用 eps 或 F.normalize A_norm F.normalize(A, p2, dim1, eps1e-8) # 处理零向量3. 内存消耗# 大矩阵的相似度矩阵可能很大 # A: (10000, 768), B: (10000, 768) # sim: (10000, 10000) 100M 元素 400MB (float32) # 使用分批计算 sim CosineSimilarity.compute_batched(A, B, batch_size1000)4. 数据类型一致性A torch.randn(5, 128, dtypetorch.float32) B torch.randn(3, 128, dtypetorch.float64) # 错误: 类型不匹配 # sim A B.T # RuntimeError # 正确: 统一类型 B B.to(A.dtype) sim CosineSimilarity.compute(A, B)5. GPU 上的效率# 在 GPU 上矩阵乘法非常高效 if torch.cuda.is_available(): A torch.randn(10000, 768, devicecuda) B torch.randn(10000, 768, devicecuda) # GPU 上计算 sim CosineSimilarity.compute(A, B) # 非常快 # 注意: 不要在 CPU 和 GPU 之间频繁切换6. 梯度传播# 余弦相似度可以传播梯度 A torch.randn(5, 128, requires_gradTrue) B torch.randn(3, 128) # 不需要梯度 sim CosineSimilarity.compute(A, B) loss -sim.diag().mean() # 对角线为正样本对 loss.backward() print(fA 的梯度: {A.grad is not None}) # True7. 半精度计算# 使用 float16 减少内存和加速 A torch.randn(10000, 768, dtypetorch.float16) B torch.randn(5000, 768, dtypetorch.float16) sim CosineSimilarity.compute(A, B) # 注意: float16 可能有精度损失总结计算两个矩阵之间所有行的余弦相似度是深度学习中的常见操作使用矩阵乘法可以高效实现。核心要点总结使用矩阵乘法F.normalize(A) F.normalize(B).T是最高效的全交叉余弦相似度计算方法避免了循环。F.cosine_similarity 是逐行配对默认行为是计算对应行的相似度不是全交叉。全交叉需要使用广播或矩阵乘法。数值稳定性使用F.normalize(eps1e-8)或添加 epsilon 避免除零问题处理全零向量。分批计算对于大规模矩阵使用compute_batched分批计算避免内存溢出。Top-K 搜索使用sim.topk(k, dim1)快速找到最相似的向量。对比学习在 SimCLR 等对比学习中使用带温度系数的余弦相似度sim / temperature。GPU 加速在 GPU 上矩阵乘法高度优化可以高效处理大规模矩阵。性能对比矩阵乘法比循环快数百倍应始终优先使用矩阵运算。通过掌握这些技巧你可以高效地计算余弦相似度矩阵应用于推荐系统、语义搜索、对比学习等各种场景。