PySpark如何实现Pandas按分类分组显示全类别聚合效果?
在PySpark中实现分组保留所有分类(含无数据类别)
要实现和Pandas中pd.cut分组后保留所有区间(包括无数据区间)的效果,核心思路是先构造包含所有目标区间的基准数据集,再与原始数据的聚合结果做左连接,具体步骤如下:
步骤1:准备数据与定义区间
首先创建PySpark DataFrame,并定义分箱区间:
from pyspark.sql import SparkSession from pyspark.sql import functions as F from pyspark.sql.types import StringType, DoubleType # 初始化SparkSession spark = SparkSession.builder.appName("GroupWithAllCategories").getOrCreate() # 创建原始数据 data = [12.0, 20.0, 40.0, 60.0, 72.0] df = spark.createDataFrame(data, DoubleType()).toDF("Age") # 定义分箱区间,和Pandas示例一致 bins = [0, 11, 30, 60] # 生成所有区间的字符串表示,匹配Pandas的cut结果格式 bin_labels = [f"({bins[i]}, {bins[i+1]}]" for i in range(len(bins)-1)]
步骤2:构造全量区间基准表
创建包含所有区间的DataFrame,作为后续左连接的基准,确保所有类别都被保留:
all_bins_df = spark.createDataFrame(bin_labels, StringType()).toDF("Age_bin")
步骤3:对原始数据分箱并聚合
使用条件判断对原始数据分箱,然后计算每个区间的均值和出现次数:
# 对Age列分箱,匹配定义的区间范围 df_with_bins = df.withColumn( "Age_bin", F.when((F.col("Age") > 0) & (F.col("Age") <= 11), "(0, 11]") .when((F.col("Age") > 11) & (F.col("Age") <= 30), "(11, 30]") .when((F.col("Age") > 30) & (F.col("Age") <= 60), "(30, 60]") # 超出最大区间的数值不纳入统计,和Pandas示例逻辑保持一致 ) # 按分箱分组,计算均值和出现次数 agg_df = df_with_bins.groupBy("Age_bin")\ .agg( F.mean("Age").alias("mean"), F.count("Age").alias("occurrences") )
步骤4:左连接保留所有区间
将全量区间表与聚合结果左连接,对无数据的区间填充空值和0:
# 左连接确保所有区间都被保留,无数据的区间填充对应值 result_df = all_bins_df.join(agg_df, on="Age_bin", how="left")\ .fillna({"mean": None, "occurrences": 0}) # 查看最终结果 result_df.show()
执行后输出结果:
+----------+----+-----------+ | Age_bin|mean|occurrences| +----------+----+-----------+ | (0, 11] |null| 0| | (11, 30] |16.0| 2| | (30, 60] |50.0| 2| +----------+----+-----------+
补充:用Bucketizer简化多区间分箱
如果区间较多,手动写when语句繁琐,可以用Bucketizer工具类简化分箱逻辑:
from pyspark.ml.feature import Bucketizer # Bucketizer为左闭右开,调整边界以匹配Pandas的右闭格式 bucketizer = Bucketizer( splits=[0.0, 11.0, 30.0, 60.0, float("inf")], inputCol="Age", outputCol="bucket_idx" ) # 生成分箱索引 df_bucketed = bucketizer.transform(df) # 将索引映射为区间字符串 bin_mapping = {i: bin_labels[i] for i in range(len(bin_labels))} df_with_bins = df_bucketed.withColumn( "Age_bin", F.map_from_entries(F.create_map(*[F.lit(k), F.lit(v)] for k, v in bin_mapping.items()))\ .getItem(F.col("bucket_idx")) ) # 后续聚合和左连接步骤与之前一致
内容的提问来源于stack exchange,提问作者Dani
相关产品推荐
相关产品推荐

