如何在Spark中实现基于cosine similarity的近似N近邻连接?
基于Spark实现余弦相似度的近似Top N近邻连接
核心思路:利用归一化向量的欧氏距离等价性
Spark ML原生没有直接支持余弦相似度的LSH实现,但可以通过归一化embedding向量+BucketedRandomProjectionLSH间接实现,因为:
- 对两个单位向量来说,余弦相似度 = 1 - (欧氏距离²)/2
- 两者的近邻排序完全一致,因此用欧氏距离的LSH找到的近邻就是余弦相似度的近邻
实现步骤(原生Spark ML方案)
1. 导入依赖并归一化embedding向量
首先使用Normalizer将A和B的embedding转换为L2范数为1的单位向量:
from pyspark.ml.feature import Normalizer from pyspark.sql import SparkSession from pyspark.sql.window import Window import pyspark.sql.functions as F spark = SparkSession.builder.appName("CosineLSH").getOrCreate() # 假设df_a、df_b为已加载的目标DataFrame,embedding列是DenseVector类型 normalizer = Normalizer(inputCol="embedding", outputCol="norm_embedding", p=2.0) df_a_norm = normalizer.transform(df_a).select("id", "text", "norm_embedding")\ .withColumnRenamed("id", "A_id").withColumnRenamed("text", "text_A") df_b_norm = normalizer.transform(df_b).select("id", "text", "norm_embedding")\ .withColumnRenamed("id", "B_id").withColumnRenamed("text", "text_B")
2. 训练BucketedRandomProjectionLSH模型
用B的归一化向量训练LSH模型,参数numHashTables控制近似精度,值越高精度越高但计算成本也越高:
from pyspark.ml.feature import BucketedRandomProjectionLSH brp = BucketedRandomProjectionLSH(inputCol="norm_embedding", outputCol="hashes", bucketLength=1.0, numHashTables=5) model = brp.fit(df_b_norm)
3. 执行近似近邻搜索并转换为余弦距离
使用approxNearestNeighbors为A中每行找到B的Top N近邻,同时支持过滤最小相似度:
N = 3 # 目标Top N数量 min_cosine_similarity = 0.7 # 最小相似度阈值 # 为A中每个向量匹配B的Top N近邻 neighbors = model.approxNearestNeighbors(df_b_norm, df_a_norm.select("norm_embedding"), N) # 关联原始字段、计算余弦距离并过滤排序 result = df_a_norm.join(neighbors, on="norm_embedding")\ .withColumn("cosine_similarity", 1 - (F.col("distance")**2)/2)\ .withColumn("cosine_distance", 1 - F.col("cosine_similarity"))\ .filter(F.col("cosine_similarity") >= min_cosine_similarity)\ .withColumn("rank", F.row_number().over(Window.partitionBy("A_id").orderBy(F.col("cosine_similarity").desc())))\ .select("A_id", "B_id", "text_A", "text_B", "rank", "cosine_distance")\ .orderBy("A_id", "rank") result.show()
如果需要批量执行相似度连接,可改用approxSimilarityJoin:
# 转换最小相似度为欧氏距离阈值:sqrt(2*(1-最小相似度)) distance_threshold = (2*(1 - min_cosine_similarity))**0.5 joined = model.approxSimilarityJoin(df_a_norm, df_b_norm, distance_threshold, distCol="euclidean_distance") # 处理结果并筛选Top N result = joined.select( F.col("datasetA.A_id").alias("A_id"), F.col("datasetB.B_id").alias("B_id"), F.col("datasetA.text_A").alias("text_A"), F.col("datasetB.text_B").alias("text_B"), (1 - (F.col("euclidean_distance")**2)/2).alias("cosine_similarity"), ((F.col("euclidean_distance")**2)/2).alias("cosine_distance") ).filter(F.col("cosine_similarity") >= min_cosine_similarity)\ .withColumn("rank", F.row_number().over(Window.partitionBy("A_id").orderBy(F.col("cosine_similarity").desc())))\ .filter(F.col("rank") <= N)\ .select("A_id", "B_id", "text_A", "text_B", "rank", "cosine_distance")\ .orderBy("A_id", "rank")
进阶方案:使用第三方近似近邻库(Spark 3.x+)
如果需要更高效率的余弦相似度近邻搜索,可使用专门的第三方库:
- spark-annoy:基于Annoy库,原生支持余弦相似度
- faiss-spark:基于Faiss库,适合高维向量的大规模搜索
spark-annoy示例
# 先安装依赖:pip install spark-annoy from spark_annoy import AnnoyIndexer # 构建Annoy索引,metric="angular"等价于余弦相似度 annoy_indexer = AnnoyIndexer( inputCol="norm_embedding", outputCol="neighbors", numTrees=10, # 树数量越多,精度越高但速度越慢 metric="angular" ) # 用B数据集构建索引 indexed_b = annoy_indexer.fit(df_b_norm) # 为A数据集搜索Top N近邻 result = indexed_b.transform(df_a_norm)\ .select("A_id", "text_A", F.explode("neighbors").alias("neighbor"))\ .select( "A_id", "text_A", F.col("neighbor.B_id").alias("B_id"), F.col("neighbor.text_B").alias("text_B"), F.col("neighbor.distance").alias("angular_distance") )\ .withColumn("cosine_similarity", 1 - F.col("angular_distance")/2)\ .withColumn("cosine_distance", F.col("angular_distance")/2)\ .filter(F.col("cosine_similarity") >= min_cosine_similarity)\ .withColumn("rank", F.row_number().over(Window.partitionBy("A_id").orderBy(F.col("cosine_similarity").desc())))\ .filter(F.col("rank") <= N)\ .select("A_id", "B_id", "text_A", "text_B", "rank", "cosine_distance")
关键注意事项
- 归一化是核心:必须确保embedding转换为单位向量,否则欧氏距离与余弦相似度的等价性不成立
- 参数调优:
numHashTables(BucketedRandomProjectionLSH)或numTrees(Annoy)需要在精度和性能之间做权衡 - 阈值转换:使用距离阈值时,务必在欧氏距离/角度距离与余弦相似度之间做好对应转换
内容的提问来源于stack exchange,提问作者Francesco Pasa
相关产品推荐
相关产品推荐

