在PySpark中为每行计算数组交集并生成最大交集列
解决Spark中计算用户列表最大交集的问题
给定如下分组后的Spark DataFrame:
import pandas as pd from pyspark.sql import SparkSession import pyspark.sql.functions as F spark = SparkSession.builder.getOrCreate() df_example = pd.DataFrame({'user1': ['u1', 'u1', 'u1', 'u5', 'u5', 'u5', 'u7','u7','u6'], 'user2': ['u2', 'u3', 'u4', 'u2', 'u4','u6','u8','u3','u6']}) sdf = spark.createDataFrame(df_example) userreposts_gr = sdf.groupby('user1').agg(F.collect_list('user2').alias('all_user2')) userreposts_gr.show()
输出:
+-----+------------+ |user1| all_user2| +-----+------------+ | u1|[u4, u2, u3]| | u7| [u8, u3]| | u5|[u4, u2, u6]| | u6| [u6]| +-----+------------+
需求是为每个user1计算其all_user2列表与其他user1的all_user2列表的交集,生成新列记录拥有最大交集的用户及交集数量,最终输出如下:
+-----+------------+------------------------------+ |user1|all_user2 |new_col | +-----+------------+------------------------------+ |u1 |[u2, u3, u4]|{max_count -> 2, user -> 'u5'}| |u5 |[u2, u4, u6]|{max_count -> 2, user -> 'u1'}| |u7 |[u8, u3] |{max_count -> 1, user -> 'u1'}| |u6 |[u6] |{max_count -> 1, user -> 'u5'}| +-----+------------+------------------------------+
解决方案步骤:
1. 自连接生成用户配对
将分组后的DataFrame与自身做自连接,排除用户与自身配对的情况,得到所有用户间的有效组合:
from pyspark.sql.window import Window # 自连接,过滤掉自身配对 pair_df = userreposts_gr.alias("a").join( userreposts_gr.alias("b"), F.col("a.user1") != F.col("b.user1"), how="inner" )
2. 计算两两用户列表的交集数量
使用array_intersect获取两个列表的交集,再用size函数统计交集元素个数:
pair_with_intersect = pair_df.withColumn( "intersect_count", F.size(F.array_intersect(F.col("a.all_user2"), F.col("b.all_user2"))) ).select( F.col("a.user1").alias("user1"), F.col("a.all_user2").alias("all_user2"), F.col("b.user1").alias("other_user"), F.col("intersect_count") )
3. 筛选每个用户的最大交集记录
通过窗口函数按user1分组,对交集数量降序排序,取排名第一的记录(即交集最大的用户):
window_spec = Window.partitionBy("user1").orderBy(F.col("intersect_count").desc()) max_intersect_df = pair_with_intersect.withColumn( "rank", F.row_number().over(window_spec) ).filter(F.col("rank") == 1).drop("rank")
4. 构造目标结果列
使用create_map函数将交集数量和对应用户组合成指定格式的Map列:
final_df = max_intersect_df.withColumn( "new_col", F.create_map( F.lit("max_count"), F.col("intersect_count"), F.lit("user"), F.col("other_user") ) ).select("user1", "all_user2", "new_col") # 查看最终结果 final_df.show(truncate=False)
执行上述代码后,即可得到符合预期的输出结果。
内容的提问来源于stack exchange,提问作者Rory
相关产品推荐
相关产品推荐

