PySpark中grouping_id函数抛出异常问题求助
PySpark grouping_id函数正确用法及问题修复
问题重现
原始代码
from pyspark.sql import SparkSession from pyspark.sql.functions import grouping_id import pyspark.sql.functions as functions spark = SparkSession.builder.appName("grouping_id_example").getOrCreate() data = [("John", "M", 25, "London", 5000), ("Jane", "F", 30, "Paris", 6000), ("Jake", "M", 45, "London", 9000), ("Julie", "F", 50, "Berlin", 10000)] df = spark.createDataFrame(data, ["Name", "Gender", "Age", "City", "Salary"]) grouped_df = df.groupBy("City", "Gender") grouped_df = grouped_df.agg({"Salary": "sum"}) grouped_df = grouped_df.withColumn("grouping_id", functions.grouping_id(["City", "Gender"]))
遇到的异常
- 传入列表时的错误:
TypeError: Invalid argument, not a string or column: ['City', 'Gender'] of type <class 'list'>. For column literals, use 'lit', 'array', 'struct' or 'create_map' function.
- 改为字符串后的错误:
pyspark.sql.utils.AnalysisException: grouping_id() can only be used with GroupingSets/Cube/Rollup;;
问题原因与解决方案
核心问题
grouping_id的作用是标识分组集合中哪些列参与了当前聚合计算,它必须配合rollup、cube或groupingSets使用,且不需要手动传入列名——函数会自动基于当前分组维度生成ID。
正确使用方式
1. 配合Rollup使用
Rollup会生成从最细粒度到全局的层级聚合结果(比如先按City+Gender聚合,再按City聚合,最后全局聚合),grouping_id用二进制位表示列是否被聚合:
- 每一位对应一个分组列,顺序与分组时的列顺序一致
- 0表示该列参与了当前分组,1表示该列被聚合(未参与分组)
示例代码:
from pyspark.sql import SparkSession from pyspark.sql.functions import grouping_id, sum spark = SparkSession.builder.appName("grouping_id_example").getOrCreate() data = [("John", "M", 25, "London", 5000), ("Jane", "F", 30, "Paris", 6000), ("Jake", "M", 45, "London", 9000), ("Julie", "F", 50, "Berlin", 10000)] df = spark.createDataFrame(data, ["Name", "Gender", "Age", "City", "Salary"]) # 使用rollup生成多层聚合结果 result_df = df.rollup("City", "Gender") \ .agg(sum("Salary").alias("total_salary")) \ .withColumn("grouping_id", grouping_id()) result_df.show()
输出结果:
+-------+------+------------+-----------+ | City|Gender|total_salary|grouping_id| +-------+------+------------+-----------+ |London | M| 14000| 0| | Paris| F| 6000| 0| | Berlin| F| 10000| 0| |London | null| 14000| 1| | Paris| null| 6000| 1| | Berlin| null| 10000| 1| | null| null| 30000| 3| +-------+------+------------+-----------+
- grouping_id=0:City和Gender都参与分组(最细粒度)
- grouping_id=1:二进制
01,表示Gender被聚合,仅City参与分组 - grouping_id=3:二进制
11,表示City和Gender都被聚合(全局聚合)
2. 配合Cube使用
Cube会生成所有可能的维度组合(比如City+Gender、City、Gender、全局),grouping_id同样用二进制位标识:
result_df = df.cube("City", "Gender") \ .agg(sum("Salary").alias("total_salary")) \ .withColumn("grouping_id", grouping_id()) result_df.show()
输出会多一行Gender单独分组的结果(grouping_id=2,二进制10,表示City被聚合):
+-------+------+------------+-----------+ | City|Gender|total_salary|grouping_id| +-------+------+------------+-----------+ |London | M| 14000| 0| | Paris| F| 6000| 0| | Berlin| F| 10000| 0| |London | null| 14000| 1| | Paris| null| 6000| 1| | Berlin| null| 10000| 1| | null| M| 14000| 2| | null| F| 16000| 2| | null| null| 30000| 3| +-------+------+------------+-----------+
3. 配合GroupingSets使用
如果需要自定义分组组合(比如只按City+Gender、全局聚合),可以用groupingSets:
from pyspark.sql import functions as F result_df = df.groupBy(F.groupingSets(["City", "Gender"], [])) \ .agg(sum("Salary").alias("total_salary")) \ .withColumn("grouping_id", grouping_id()) result_df.show()
输出:
+-------+------+------------+-----------+ | City|Gender|total_salary|grouping_id| +-------+------+------------+-----------+ |London | M| 14000| 0| | Paris| F| 6000| 0| | Berlin| F| 10000| 0| | null| null| 30000| 3| +-------+------+------------+-----------+
总结
grouping_id无需传入列名,自动基于当前分组集合的维度生成ID- 必须配合
rollup、cube或groupingSets使用,普通groupBy无法调用该函数 - grouping_id的数值是二进制转十进制的结果,每一位对应分组列是否参与聚合(0=参与,1=聚合)
内容的提问来源于stack exchange,提问作者Sachin Sukumaran
相关产品推荐
相关产品推荐

