MS Fabric Spark Job计算余弦相似度矩阵遇性能瓶颈求助
解决PySpark计算商品余弦相似度时CrossJoin停滞的问题
问题背景
现有40k条去重后的商品数据,需计算每个商品的Top5相似商品。用sklearn在Google Colab TPU上10分钟可完成,但在Microsoft Fabric Spark Job中执行CrossJoin操作时直接停滞,无报错日志。
核心原因
你的PySpark代码陷入停滞并非平台资源限制,而是逻辑设计错误:
- 40k条数据做CrossJoin会生成
40000 * 40000 = 1.6e9行数据,远超Spark集群的存储和计算能力,哪怕是大内存节点也无法处理这么大规模的中间数据。 - sklearn的
cosine_similarity底层依赖高效的BLAS/LAPACK矩阵运算,且np.argpartition仅计算TopN索引,无需生成全量相似度矩阵;而你原PySpark代码试图生成所有两两组合,属于典型的O(n²)低效实现。
优化方案
改用Spark MLlib的BucketedRandomProjectionLSH算法,专门针对大规模高维数据做近似最近邻检索,无需生成全量笛卡尔积,大幅降低计算和存储开销:
- 预计算每个商品向量的L2范数,避免重复计算
- 使用LSH生成候选相似商品集(而非全量配对)
- 对候选集计算精确余弦相似度
- 用窗口函数提取每个商品的Top5相似项
完整优化代码
from pyspark.ml.feature import Tokenizer, StopWordsRemover, CountVectorizer from pyspark.ml.feature import BucketedRandomProjectionLSH from pyspark.sql import functions as F from pyspark.sql.window import Window # 数据预处理(和原代码一致) articulos = articulos.select("article_id", "product_code", "detail_desc") articulos = articulos.dropDuplicates(["product_code"]).dropna(subset=["detail_desc"]) tokenizer = Tokenizer(inputCol="detail_desc", outputCol="words") articulos_tokenized = tokenizer.transform(articulos) remover = StopWordsRemover(inputCol="words", outputCol="filtered_words") articulos_clean = remover.transform(articulos_tokenized) vectorizer = CountVectorizer(inputCol="filtered_words", outputCol="features") vectorizer_model = vectorizer.fit(articulos_clean) articulos_final = vectorizer_model.transform(articulos_clean) # 预计算每个向量的L2范数 articulos_final = articulos_final.withColumn( "norm", F.expr("features.norm(2)") ).select("article_id", "features", "norm") # 初始化LSH模型 lsh = BucketedRandomProjectionLSH( inputCol="features", outputCol="hashes", bucketLength=1.0, numHashTables=5 # 哈希表数量,值越高召回率越高,计算量越大 ) lsh_model = lsh.fit(articulos_final) # 用LSH检索候选相似商品(每个商品返回近似相似的候选集) candidates = lsh_model.approxSimilarityJoin( articulos_final, articulos_final, threshold=2.0, # 余弦相似度对应的L2距离阈值,可调整 distCol="l2_distance" ).select( F.col("datasetA.article_id").alias("article_id"), F.col("datasetB.article_id").alias("similar_article_id"), F.col("datasetA.features").alias("features_a"), F.col("datasetB.features").alias("features_b"), F.col("datasetA.norm").alias("norm_a"), F.col("datasetB.norm").alias("norm_b") ) # 计算精确余弦相似度,排除自身匹配 candidates = candidates.filter(F.col("article_id") != F.col("similar_article_id")).withColumn( "cosine_similarity", F.col("features_a").dot(F.col("features_b")) / (F.col("norm_a") * F.col("norm_b")) ).drop("features_a", "features_b", "norm_a", "norm_b") # 提取每个商品的Top5相似商品 window_spec = Window.partitionBy("article_id").orderBy(F.col("cosine_similarity").desc()) top_5_simart = candidates.withColumn( "rank", F.row_number().over(window_spec) ).filter(F.col("rank") <= 5).drop("rank") # 保存结果 top_5_simart.write.mode('overwrite').json(Top5SimartPath)
额外优化建议
- 调整LSH参数:
bucketLength和numHashTables需要根据你的数据特征调整,numHashTables越大召回率越高,但计算量也会增加;threshold对应余弦相似度的阈值,可根据需求调整。 - 资源配置:在Microsoft Fabric中,可调整Spark集群的节点规格(如选用更大的VM实例),增加Executor内存和核心数,提升分布式计算效率。
- 向量优化:可考虑用TF-IDF替代CountVectorizer,提升文本特征的代表性,进而提升相似度计算的准确性。
内容的提问来源于stack exchange,提问作者JLd
相关产品推荐
相关产品推荐

