You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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算法,专门针对大规模高维数据做近似最近邻检索,无需生成全量笛卡尔积,大幅降低计算和存储开销:

  1. 预计算每个商品向量的L2范数,避免重复计算
  2. 使用LSH生成候选相似商品集(而非全量配对)
  3. 对候选集计算精确余弦相似度
  4. 用窗口函数提取每个商品的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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.17 12:57:11