如何优化这段Numpy余弦距离计算代码以提升运行速度?
优化余弦相似度计算的方案
原代码的性能瓶颈集中在重复解析字符串和逐行循环计算上,针对数万条记录的场景,以下是针对性的加速方案:
核心优化思路
- 一次性完成所有embedding字符串的解析,避免重复计算
- 提前预计算所有embedding的L2范数,减少重复运算
- 利用numpy矩阵批量运算替代Python循环,充分发挥底层优化能力
完整优化代码
import pandas as pd import numpy as np import ast # 读取原始数据 df = pd.read_csv("https://testme162.s3.amazonaws.com/cosign_dist.csv") # 1. 一次性解析所有embedding字符串为numpy数组 # 列表推导式比apply方法的解析效率更高 embeddings = np.array([np.array(ast.literal_eval(s)) for s in df["embeddings"]]) # 2. 提前计算所有embedding的L2范数,keepdims保证维度匹配 norms = np.linalg.norm(embeddings, axis=1, keepdims=True) # 3. 矩阵批量计算余弦相似度 # 矩阵乘法实现两两内积,再除以范数的外积得到相似度矩阵 cos_sim_matrix = embeddings @ embeddings.T / (norms @ norms.T) # 将相似度矩阵转换为原代码格式的列,并与原DataFrame合并 dist_columns = pd.DataFrame(cos_sim_matrix, columns=[f"dist{k}" for k in range(len(df))]) final_df = pd.concat([df, dist_columns], axis=1)
性能提升说明
- 解析效率:原代码每次循环都会重复解析所有embedding字符串,总解析次数为O(N²);优化后仅解析一次,次数为O(N),数万条记录下可减少百万级的重复操作
- 计算效率:numpy的矩阵运算基于BLAS/LAPACK底层优化,比Python循环+apply快数倍到数十倍
- 范数计算:提前预计算所有范数,避免O(N²)次重复的范数求解运算
内容的提问来源于stack exchange,提问作者shantanuo
相关产品推荐
相关产品推荐

