如何在PySpark中实现SQL ROLLUP逻辑,添加州级总计行?
Got it, let's tackle this! You want to replicate the SQL ROLLUP behavior in PySpark where each state's data is followed by its total row. Here's a step-by-step solution that matches your desired output:
Step 1: Set Up Sample DataFrame
First, let's define the input data to work with (matching your example):
from pyspark.sql import SparkSession from pyspark.sql.functions import sum, when, col # Initialize Spark session spark = SparkSession.builder.appName("StateTotals").getOrCreate() # Sample input data data = [ ("WA", "Seattle", 10), ("WA", "Tacoma", 11), ("MA", "Boston", 11), ("MA", "Cambridge", 3), ("MA", "Quincy", 5) ] df = spark.createDataFrame(data, ["state", "city", "count"])
Step 2: Use ROLLUP to Calculate Aggregates
PySpark's rollup function lets you compute hierarchical aggregates, just like SQL ROLLUP. We'll group by state and city to get both individual city counts and state-level totals:
# Compute rollup aggregates (includes city-level counts and state totals) rollup_df = df.rollup("state", "city").agg(sum("count").alias("count"))
This generates rows where city is null for state totals, plus an extra grand total row (with both state and city null) which we'll filter out next.
Step 3: Clean Up Rows
We don't need the grand total, so let's remove that and replace null city values with "Total" for clarity:
# Filter out the grand total row (where state is null) filtered_df = rollup_df.filter(col("state").isNotNull()) # Replace null city values with "Total" for state total rows result_df = filtered_df.withColumn("city", when(col("city").isNull(), "Total").otherwise(col("city")))
Step 4: Order Results Correctly
To ensure the total row comes after all cities in each state, add a temporary sort key and order the data:
# Add a sort key: 0 for city rows, 1 for total rows result_df = result_df.withColumn("sort_order", when(col("city") == "Total", 1).otherwise(0)) # Order by state, then sort order (cities first), then city name result_df = result_df.orderBy("state", "sort_order", "city")
Step 5: View the Final Output
Drop the temporary sort_order column and show the result:
result_df.drop("sort_order").show()
Final Output:
+-----+---------+-----+ |state| city|count| +-----+---------+-----+ | MA| Boston| 11| | MA|Cambridge| 3| | MA| Quincy| 5| | MA| Total| 19| | WA| Seattle| 10| | WA| Tacoma| 11| | WA| Total| 21| +-----+---------+-----+
This exactly matches your desired output—each state's cities are listed first, followed by the state's total row.
内容的提问来源于stack exchange,提问作者yokielove

