PySpark实现子区域映射大区并按Streak统计行数
PySpark实现子区域到大区映射及streak统计的方案
一、核心思路拆解
你可以分四步完成需求:先构建大区与子区域的映射表,再统计子区域的streak指标,接着关联映射关系,最后按需聚合到大区层级。
二、具体实现步骤
1. 构建大区-子区域映射DataFrame
把给定的映射关系转换成Spark能识别的DataFrame,后续通过关联操作完成子区域到大区的映射,比写一堆when条件更易维护:
from pyspark.sql import SparkSession from pyspark.sql import functions as F # 初始化SparkSession(如果还没初始化) spark = SparkSession.builder.appName("RegionStats").getOrCreate() # 定义大区-子区域映射数据 region_mapping_data = [ ("APAC", "Southern Asia"), ("APAC", "South-Eastern Asia"), ("APAC", "Central Asia"), ("EU/UK", "Western Europe"), ("EU/UK", "Eastern Europe"), ("EU/UK", "Northern Europe"), ("EU/UK", "Southern Europe"), ("MEA", "Western Asia"), ("MEA", "Western Africa"), ("MEA", "Southern Africa"), ("MEA", "Northern Africa"), ("MEA", "Eastern Asia"), ("MEA", "Eastern Africa"), ("NA", "North America"), ("LATAM", "Caribbean"), ("LATAM", "Central America"), ("LATAM", "South America") ] # 创建映射DataFrame region_mapping_df = spark.createDataFrame(region_mapping_data, ["region", "subregion"])
2. 统计每个子区域的streak指标
用groupBy+sum(when(...))组合统计满足条件的行数:
# 统计每个subregion的streak==1和streak>=3的行数 subregion_stats_df = df.groupBy("subregion")\ .agg( # 统计streak=1的行数 F.sum(F.when(F.col("streak") == 1, 1).otherwise(0)).alias("streak_1_count"), # 统计streak>=3的行数 F.sum(F.when(F.col("streak") >= 3, 1).otherwise(0)).alias("streak_ge3_count") )
3. 关联映射表与子区域统计结果
通过join操作把每个子区域对应到大区:
# 关联子区域统计结果和映射表 joined_df = subregion_stats_df.join(region_mapping_df, on="subregion", how="inner")
- 如果你担心有未在映射表中的子区域,可以把
how="inner"改成how="left",后续用F.coalesce(F.col("region"), F.lit("Other"))把未映射的子区域归类到"Other"大区。
4. (可选)按大区汇总统计结果
如果需要最终结果按大区聚合,再做一次groupBy求和:
# 按大区汇总统计数据 region_summary_df = joined_df.groupBy("region")\ .agg( F.sum("streak_1_count").alias("total_streak_1"), F.sum("streak_ge3_count").alias("total_streak_ge3") )
三、完整代码示例
from pyspark.sql import SparkSession from pyspark.sql import functions as F # 初始化SparkSession spark = SparkSession.builder.appName("RegionStats").getOrCreate() # 1. 构建大区-子区域映射表 region_mapping_data = [ ("APAC", "Southern Asia"), ("APAC", "South-Eastern Asia"), ("APAC", "Central Asia"), ("EU/UK", "Western Europe"), ("EU/UK", "Eastern Europe"), ("EU/UK", "Northern Europe"), ("EU/UK", "Southern Europe"), ("MEA", "Western Asia"), ("MEA", "Western Africa"), ("MEA", "Southern Africa"), ("MEA", "Northern Africa"), ("MEA", "Eastern Asia"), ("MEA", "Eastern Africa"), ("NA", "North America"), ("LATAM", "Caribbean"), ("LATAM", "Central America"), ("LATAM", "South America") ] region_mapping_df = spark.createDataFrame(region_mapping_data, ["region", "subregion"]) # 2. 统计子区域的streak指标 subregion_stats_df = df.groupBy("subregion")\ .agg( F.sum(F.when(F.col("streak") == 1, 1).otherwise(0)).alias("streak_1_count"), F.sum(F.when(F.col("streak") >= 3, 1).otherwise(0)).alias("streak_ge3_count") ) # 3. 关联映射表 joined_df = subregion_stats_df.join(region_mapping_df, on="subregion", how="inner") # 4. 按大区汇总(可选) region_summary_df = joined_df.groupBy("region")\ .agg( F.sum("streak_1_count").alias("total_streak_1"), F.sum("streak_ge3_count").alias("total_streak_ge3") ) # 查看结果 region_summary_df.show()
内容的提问来源于stack exchange,提问作者Strayhorn
相关产品推荐
相关产品推荐

