PySpark实现按Home_City统计去重跨城访问人数矩阵
高效构建PySpark访问城市统计矩阵解决方案
数据集结构
Person_Info表
| Person_ID | Home_City |
|---|---|
| 1 | New_York |
| 2 | New_York |
| 3 | Miami |
| 4 | Chicago |
Visit_Info表
| Person_ID | Visit_City |
|---|---|
| 1 | New_York |
| 2 | Miami |
| 2 | Miami |
| 2 | Miami |
| 3 | Miami |
| 3 | Chicago |
| 4 | Chicago |
需求说明
按Home_City统计两类数据:
- 该城市的总人数(
Total_People) - 该城市人群去过的不同访问城市的人数(同一人多次访问同一城市仅计1次)
最终输出矩阵格式如下:
| Home_City | Total_People | New_York | Miami | Chicago |
|---|---|---|---|---|
| New_York | 2 | 1 | 1 | 0 |
| Miami | 1 | 0 | 1 | 1 |
| Chicago | 1 | 0 | 0 | 1 |
原方案问题
原代码仅实现了前两列,但后续用Python循环扩展列时效率极低,还出现java.lang.StackOverflowError——这是因为循环会触发多次Spark作业,且collect_list会把大量数据拉到Driver端,超出内存承载上限。
高效解决方案
利用Spark原生的去重(distinct)、分组聚合(groupBy)、**透视(pivot)**算子实现全程分布式计算,彻底规避Driver端压力和低效循环:
步骤说明
- 清洗访问数据:对
Visit_Info按Person_ID+Visit_City去重,确保同一人同一城市仅保留一条有效记录 - 关联家乡信息:将清洗后的访问数据与
Person_Info关联,匹配每个人的家乡城市 - 分组统计+透视:按
Home_City分组,要么固定指定访问城市列统计人数,要么用pivot动态生成列,同时计算各家乡城市的总人数 - 填充空值:将透视后缺失的列值填充为0,保证矩阵完整性
完整代码
from pyspark.sql import functions as F # -------------------------- # 方案1:已知所有访问城市的固定列实现 # -------------------------- # 1. 清洗访问数据,去重同一人同一城市的重复记录 clean_visit = df_visit_info.distinct() # 2. 关联家乡信息 home_visit_join = clean_visit.join(df_person_info, on="Person_ID", how="inner") # 3. 分组统计总人数+各访问城市的人数 result_fixed = home_visit_join.groupBy("Home_City") \ .agg( # 统计家乡城市总人数(去重后的Person_ID数量) F.countDistinct("Person_ID").alias("Total_People"), # 统计去过New_York的人数 F.countDistinct(F.when(F.col("Visit_City") == "New_York", F.col("Person_ID"))).alias("New_York"), # 统计去过Miami的人数 F.countDistinct(F.when(F.col("Visit_City") == "Miami", F.col("Person_ID"))).alias("Miami"), # 统计去过Chicago的人数 F.countDistinct(F.when(F.col("Visit_City") == "Chicago", F.col("Person_ID"))).alias("Chicago") ) \ # 4. 空值填充为0(没人去过的城市会返回null) .fillna(0, subset=["New_York", "Miami", "Chicago"]) # -------------------------- # 方案2:访问城市未知的动态列实现(更灵活) # -------------------------- # 获取所有唯一的访问城市 visit_cities = [row["Visit_City"] for row in clean_visit.select("Visit_City").distinct().collect()] # 用pivot动态生成访问城市列,统计各城市的访问人数 result_dynamic = home_visit_join.groupBy("Home_City") \ .pivot("Visit_City", visit_cities) \ .agg(F.countDistinct("Person_ID")) \ # 关联总人数统计结果 .join( df_person_info.groupBy("Home_City").agg(F.countDistinct("Person_ID").alias("Total_People")), on="Home_City", how="inner" ) \ # 调整列顺序,将Total_People放在第二列 .select("Home_City", "Total_People", *visit_cities) \ # 空值填充为0 .fillna(0, subset=visit_cities) # 展示结果 result_fixed.show() result_dynamic.show()
方案优势
- 分布式执行:所有计算在Spark集群上完成,避免将数据拉到Driver端,彻底解决栈溢出问题
- 高效低耗:利用Spark原生算子替代Python循环,减少作业开销,提升执行效率
- 灵活适配:提供固定列和动态列两种实现,适配已知/未知访问城市的业务场景
内容的提问来源于stack exchange,提问作者Daryl Clark
相关产品推荐
相关产品推荐

