Spark DataFrame按ID分组实现行号与组内总数统计
Hey there! Let's work through your Spark problem step by step. You need two key outputs for each id group: a 0-based running row number, and the total count of records in the group. Here's how to achieve this with both the Spark Scala API and Spark SQL:
Using Spark Scala API
We'll use window functions here—they're ideal for per-group calculations like this.
First, import the required functions and define your window specifications:
import org.apache.spark.sql.expressions.Window import org.apache.spark.sql.functions.{row_number, count} // Window for grouping by id and sorting rows (pick a column to sort for consistent row numbering) val idPartitionWindow = Window.partitionBy("id").orderBy("val") // Window just for grouping by id (to calculate total group count) val idCountWindow = Window.partitionBy("id")
Then apply these windows to your DataFrame:
val resultDF = df .withColumn("order", row_number().over(idPartitionWindow) - 1) // Subtract 1 to get 0-based index .withColumn("count", count("*").over(idCountWindow)) // View the final output resultDF.show()
Quick Notes:
row_number()starts counting from 1 by default, so subtracting 1 gives you the 0-based sequence you need.- The
orderByclause is mandatory forrow_number()—it ensures consistent row ordering within each group. Swap"val"with another column if you need a different sort order. count("*").over(idCountWindow)calculates the total number of records peridgroup and attaches that value to every row in the group.
Using Spark SQL
If you prefer writing SQL queries, you can get the same result with Spark SQL:
First, register your DataFrame as a temporary view:
df.createOrReplaceTempView("records")
Then run this SQL query:
SELECT id, val, ROW_NUMBER() OVER (PARTITION BY id ORDER BY val) - 1 AS `order`, COUNT(*) OVER (PARTITION BY id) AS count FROM records
Execute the query and check the output:
val resultDF = spark.sql(""" SELECT id, val, ROW_NUMBER() OVER (PARTITION BY id ORDER BY val) - 1 AS `order`, COUNT(*) OVER (PARTITION BY id) AS count FROM records """) resultDF.show()
This will generate exactly the output you're expecting—each row gets its 0-based position in the group, plus the total number of rows in that group.
内容的提问来源于stack exchange,提问作者Brian

