如何在PySpark中构建基于用户交集/并集的物品-物品交互矩阵
PySpark 实现物品-物品交互矩阵
前置准备
首先导入所需函数、初始化SparkSession,样例数据集构造如下(实际使用时替换为你自己的DataFrame即可):
from pyspark.sql import SparkSession from pyspark.sql.functions import col, collect_set, explode, countDistinct, udf from pyspark.sql.types import IntegerType # 初始化SparkSession spark = SparkSession.builder.appName("item_item_matrix").getOrCreate() # 构造测试样例数据 data = [ ("user1", "A"), ("user1", "B"), ("user2", "A"), ("user3", "B"), ("user4", "C") ] df = spark.createDataFrame(data, schema=["user_id", "item_id"])
1、共同用户(交集)矩阵实现
逻辑:先聚合每个用户对应的所有物品,再将每个用户的物品两两配对,配对出现的次数即为两个物品的共同用户数,最后透视得到矩阵格式。
# 聚合每个用户的物品集合 user_item_df = df.groupBy("user_id").agg(collect_set("item_id").alias("items")) # 生成所有两两物品配对(包含物品自身配对) pair_df = user_item_df.select( explode("items").alias("item1"), col("items").alias("item_list") ).select( "item1", explode("item_list").alias("item2") ) # 统计每对物品的共同用户数(交集大小) intersection_count_df = pair_df.groupBy("item1", "item2").count() # 透视转为矩阵格式 intersection_matrix = intersection_count_df.groupBy("item1").pivot("item2").sum("count").fillna(0) intersection_matrix.show()
运行输出符合预期:
+-----+---+---+---+ |item1| A| B| C| +-----+---+---+---+ | A| 2| 1| 0| | B| 1| 2| 0| | C| 0| 0| 1| +-----+---+---+---+
2、用户并集矩阵实现
逻辑:用集合公式|A ∪ B| = |A| + |B| - |A ∩ B|直接基于交集结果计算,无需重复生成物品配对,性能更高。
# 先计算每个物品的独立用户数 item_user_count = df.groupBy("item_id").agg(countDistinct("user_id").alias("user_cnt")) item_cnt_map = {row["item_id"]: row["user_cnt"] for row in item_user_count.collect()} # 注册UDF计算并集大小 def calc_union(item1, item2, intersect_cnt): return item_cnt_map[item1] + item_cnt_map[item2] - intersect_cnt union_udf = udf(calc_union, IntegerType()) # 计算每对物品的并集大小 union_count_df = intersection_count_df.withColumn("union_cnt", union_udf(col("item1"), col("item2"), col("count"))) # 透视转为矩阵格式 union_matrix = union_count_df.groupBy("item1").pivot("item2").sum("union_cnt").fillna(0) union_matrix.show()
运行输出符合预期:
+-----+---+---+---+ |item1| A| B| C| +-----+---+---+---+ | A| 2| 3| 3| | B| 3| 2| 3| | C| 3| 3| 1| +-----+---+---+---+
注意事项
如果物品数量非常多,pivot操作会占用较多资源,可先过滤掉不需要的物品对再执行透视,避免生成过多无效列。
内容的提问来源于stack exchange,提问作者Adi Singh
相关产品推荐
相关产品推荐

