求矩阵间余弦距离矩阵的GPU加速实现方案
解决方案
核心思路
余弦距离的本质是 1 - 余弦相似度,而余弦相似度在向量归一化后等价于向量的点积。利用PyTorch/TensorFlow的GPU加速矩阵运算,可以高效处理大规模的三维嵌入数据。
场景1:计算每个样本内部token向量的成对余弦距离
针对你的[10663,512,768]数据(10663个样本,每个样本包含512个768维token向量),计算每个样本内部512×512的成对余弦距离矩阵:
PyTorch实现(GPU加速)
import torch # 假设你的嵌入数据为numpy数组,形状[10663,512,768] embeddings = ... # 自动切换到可用GPU device = torch.device("cuda" if torch.cuda.is_available() else "cpu") emb_tensor = torch.tensor(embeddings, dtype=torch.float32).to(device) # 对每个token向量做L2归一化(沿特征维度) norm_emb = torch.nn.functional.normalize(emb_tensor, p=2, dim=-1) # 批量矩阵乘法计算余弦相似度,结果形状[10663,512,512] cos_sim = torch.bmm(norm_emb, norm_emb.transpose(1, 2)) # 转换为余弦距离 cos_dist = 1 - cos_sim # 可选:转回numpy数组(移回CPU) cos_dist_np = cos_dist.cpu().numpy()
TensorFlow实现(GPU加速)
import tensorflow as tf # 假设你的嵌入数据为numpy数组,形状[10663,512,768] embeddings = ... emb_tensor = tf.convert_to_tensor(embeddings, dtype=tf.float32) # L2归一化 norm_emb = tf.math.l2_normalize(emb_tensor, axis=-1) # 批量矩阵乘法计算余弦相似度 cos_sim = tf.matmul(norm_emb, norm_emb, transpose_b=True) # 转换为余弦距离 cos_dist = 1 - cos_sim # 可选:转回numpy数组 cos_dist_np = cos_dist.numpy()
场景2:计算样本间的成对余弦距离
如果需要将每个样本的512个token向量聚合为一个样本级向量(比如平均池化),再计算10663×10663的样本间成对余弦距离:
import torch emb_tensor = torch.tensor(embeddings, dtype=torch.float32).to(device) # 对每个样本的token向量取平均,得到[10663,768]的样本级向量 sample_emb = torch.mean(emb_tensor, dim=1) # 归一化 norm_sample_emb = torch.nn.functional.normalize(sample_emb, p=2, dim=-1) # 计算样本间余弦相似度与距离 sample_cos_sim = torch.mm(norm_sample_emb, norm_sample_emb.T) sample_cos_dist = 1 - sample_cos_sim
内存优化(针对超大规模数据)
如果GPU内存不足,可分批次处理:
import torch batch_size = 64 # 根据GPU内存调整批次大小 cos_dist_list = [] for idx in range(0, embeddings.shape[0], batch_size): batch = torch.tensor(embeddings[idx:idx+batch_size], dtype=torch.float32).to(device) norm_batch = torch.nn.functional.normalize(batch, p=2, dim=-1) batch_cos_dist = 1 - torch.bmm(norm_batch, norm_batch.transpose(1,2)) cos_dist_list.append(batch_cos_dist.cpu()) # 合并所有批次结果 cos_dist = torch.cat(cos_dist_list, dim=0).numpy()
优势说明
- 完全利用GPU并行计算,处理10663个样本的效率远高于CPU版的
sklearn.metrics.pairwise.cosine_distances - 代码简洁,逻辑和
cosine_distances一致,无需复杂封装 - 支持批量处理,适配不同GPU内存规模
内容的提问来源于stack exchange,提问作者David
相关产品推荐
相关产品推荐

