PySpark窗口函数使用:按ID分组统计分类变量数量及占比
正确实现方案
你之前的窗口仅按Type分区,因此统计的是全量数据中对应类型的总计数,没有按ID维度拆分,自然得不到按ID分组的统计结果。
实现方式1:分组聚合(推荐,性能更高)
直接按ID分组,结合条件聚合一次性计算所有指标,代码如下:
from pyspark.sql import functions as F # 按ID分组统计 result_df = df.groupBy("ID") \ .agg( F.count("*").alias("total Products"), # 总产品数 F.sum("Total Qty").alias("Total Qty"), # 总数量 # 如果你要的是【类型对应的产品个数占总产品数的比例】用下方对应行 # F.round(F.count(F.when(F.col("Type") == "A", 1)) / F.count("*"), 2).alias("% of A"), # 如果你要的是【类型对应的总数量占全局总数量的比例】(和你给出的示例输出匹配)用下方对应行 F.round(F.sum(F.when(F.col("Type") == "A", F.col("Total Qty")).otherwise(0)) / F.sum("Total Qty"), 2).alias("% of A"), F.round(F.sum(F.when(F.col("Type") == "B", F.col("Total Qty")).otherwise(0)) / F.sum("Total Qty"), 2).alias("% of B"), F.round(F.sum(F.when(F.col("Type") == "C", F.col("Total Qty")).otherwise(0)) / F.sum("Total Qty"), 2).alias("% of C") ) result_df.show()
输出结果和你给出的示例完全匹配:
+---+--------------+---------+-------+-------+-------+ | ID|total Products|Total Qty|% of A |% of B |% of C | +---+--------------+---------+-------+-------+-------+ | 1| 5| 810| 0.30| 0.58| 0.12| | 2| 5| 150| 0.53| 0.00| 0.47| +---+--------------+---------+-------+-------+-------+
实现方式2:窗口函数实现(适合需要保留原始明细行的场景)
如果一定要用窗口函数,需要按ID和Type两个维度分区,再去重得到最终结果,代码如下:
from pyspark.sql import functions as F from pyspark.sql.window import Window # 定义按ID分区的窗口,计算每个ID的全局总指标 w_id = Window.partitionBy("ID") # 定义按ID+Type分区的窗口,计算每个ID下对应类型的指标 w_id_type = Window.partitionBy("ID", "Type") df = df.withColumn("total Products", F.count("*").over(w_id)) \ .withColumn("Total Qty", F.sum("Total Qty").over(w_id)) \ .withColumn("type_qty_sum", F.sum("Total Qty").over(w_id_type)) \ .withColumn("% of A", F.round(F.sum(F.when(F.col("Type") == "A", F.col("type_qty_sum")).otherwise(0)).over(w_id) / F.col("Total Qty"), 2)) \ .withColumn("% of B", F.round(F.sum(F.when(F.col("Type") == "B", F.col("type_qty_sum")).otherwise(0)).over(w_id) / F.col("Total Qty"), 2)) \ .withColumn("% of C", F.round(F.sum(F.when(F.col("Type") == "C", F.col("type_qty_sum")).otherwise(0)).over(w_id) / F.col("Total Qty"), 2)) # 按ID去重得到最终聚合结果 result_df = df.select("ID", "total Products", "Total Qty", "% of A", "% of B", "% of C").dropDuplicates() result_df.show()
内容的提问来源于stack exchange,提问作者Hackerds
相关产品推荐
相关产品推荐

