资源受限下,如何优化PySpark代码高效计算大规模数据集的Jaccard相似度
问题核心分析
原代码卡住的致命原因是全量交叉连接(crossJoin):假设ItemA有N个唯一值,交叉后会生成N²条记录,哪怕N是100万,这个量级也远超现有集群的处理能力,直接导致任务无法推进。此外,将DataFrame转为RDD处理会丢失Spark Catalyst的优化,进一步降低效率。
优化方案(无需增加集群资源)
核心思路是只计算有实际交集的ItemA对(Jaccard相似度非零的前提是两个ItemA共享至少一个ItemB),大幅减少计算量,同时用Spark内置函数替代RDD操作提升效率。
步骤1:计算每个ItemA的ItemB集合大小
先统计每个ItemA对应的去重ItemB数量,后续用于Jaccard公式计算:
from pyspark.sql import functions as F # 计算每个ItemA的去重ItemB数量 item_b_size = df.groupBy("ItemA").agg(F.countDistinct("ItemB").alias("itemb_count"))
步骤2:生成有共同ItemB的ItemA对并统计交集次数
通过ItemB反向关联,找到所有共享同一ItemB的ItemA对,统计它们的交集(即共同出现的ItemB数量):
# 按ItemB分组,收集对应的ItemA集合 itema_groups = df.groupBy("ItemB").agg(F.collect_set("ItemA").alias("itema_list")) # 过滤掉只对应单个ItemA的ItemB(无法生成有效对) itema_groups = itema_groups.filter(F.size(F.col("itema_list")) >= 2) # 生成每个ItemB下的ItemA两两组合(只保留i<j避免重复计算) def generate_valid_pairs(itema_list): pairs = [] list_len = len(itema_list) for i in range(list_len): for j in range(i + 1, list_len): pairs.append((itema_list[i], itema_list[j])) return pairs # 统计每个ItemA对的交集次数(即共同ItemB的数量) pair_intersection = ( itema_groups.rdd .flatMap(lambda row: generate_valid_pairs(row["itema_list"])) .map(lambda pair: ((pair[0], pair[1]), 1)) .reduceByKey(lambda a, b: a + b) .toDF(["itema_pair", "intersection_count"]) .select( F.col("itema_pair._1").alias("ItemA_i"), F.col("itema_pair._2").alias("ItemA_j"), "intersection_count" ) )
步骤3:计算Jaccard相似度
关联之前统计的ItemB数量,用公式Jaccard = 交集数 / (|A| + |B| - 交集数)计算相似度:
# 关联ItemA_i的ItemB数量 similarity_df = ( pair_intersection .join(item_b_size, F.col("ItemA_i") == F.col("ItemA"), how="left") .withColumnRenamed("itemb_count", "size_i") .drop("ItemA") # 关联ItemA_j的ItemB数量 .join(item_b_size, F.col("ItemA_j") == F.col("ItemA"), how="left") .withColumnRenamed("itemb_count", "size_j") .drop("ItemA") # 计算Jaccard相似度,处理分母为0的情况 .withColumn( "jaccard_sim", F.when( (F.col("size_i") + F.col("size_j") - F.col("intersection_count")) > 0, F.col("intersection_count") / (F.col("size_i") + F.col("size_j") - F.col("intersection_count")) ).otherwise(0.0) ) ) # 查看结果 similarity_df.show(10, truncate=False)
额外优化建议
- 编码ItemB减少内存占用:用
StringIndexer将字符串类型的ItemB转为整数,降低collect_set的内存消耗:
from pyspark.ml.feature import StringIndexer indexer = StringIndexer(inputCol="ItemB", outputCol="ItemB_idx") df_indexed = indexer.fit(df).transform(df).drop("ItemB") # 后续用ItemB_idx替代ItemB进行所有计算
- 调整分区数:根据集群核心数设置分区数(建议为核心数的2-3倍),避免分区过多或过少:
# 在分组前调整分区 df = df.repartition(200) # 替换为适合集群的数值
- 补充对称对(可选):如果需要
(ItemA_j, ItemA_i)的对称记录,可以复制并交换列:
symmetric_df = similarity_df.select( F.col("ItemA_j").alias("ItemA_i"), F.col("ItemA_i").alias("ItemA_j"), "jaccard_sim" ) final_df = similarity_df.union(symmetric_df)
内容的提问来源于stack exchange,提问作者Rayne
相关产品推荐
相关产品推荐

