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

资源受限下,如何优化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)
额外优化建议
  1. 编码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进行所有计算
  1. 调整分区数:根据集群核心数设置分区数(建议为核心数的2-3倍),避免分区过多或过少:
# 在分组前调整分区
df = df.repartition(200)  # 替换为适合集群的数值
  1. 补充对称对(可选):如果需要(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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.16 05:58:09