PySpark中为连续相同colA+colB分组添加连续序号的问题求助
连续相同colA-colB组合的序号重置问题
我正在编写PySpark代码,需要为已按colA和Date排序的DataFrame,针对连续相同的colA与colB组合分组添加连续序号。
原始DataFrame
| colA | colB | Date |
|---|---|---|
| A | 1 | 01-01-2014 |
| A | 1 | 01-02-2014 |
| A | 3 | 30-04-2014 |
| A | 3 | 05-05-2014 |
| A | 2 | 25-05-2014 |
| A | 1 | 06-06-2014 |
| A | 1 | 21-07-2014 |
| B | 1 | 04-09-2014 |
| B | 1 | 19-10-2014 |
| B | 1 | 03-12-2014 |
| C | 3 | 17-01-2015 |
| C | 2 | 03-03-2015 |
| C | 2 | 17-04-2015 |
预期结果
| colA | colB | Date | ROWNUM |
|---|---|---|---|
| A | 1 | 01-01-2014 | 1 |
| A | 1 | 01-02-2014 | 2 |
| A | 3 | 30-04-2014 | 1 |
| A | 3 | 05-05-2014 | 2 |
| A | 2 | 25-05-2014 | 1 |
| A | 1 | 06-06-2014 | 1 |
| A | 1 | 21-07-2014 | 2 |
| B | 1 | 04-09-2014 | 1 |
| B | 1 | 19-10-2014 | 2 |
| B | 1 | 03-12-2014 | 3 |
| C | 3 | 17-01-2015 | 1 |
| C | 2 | 03-03-2015 | 1 |
| C | 2 | 17-04-2015 | 2 |
错误尝试结果
使用row_number()函数时,结果不符合预期:当colA为A、colB为1的组合再次出现时,序号从3开始而非重置为1,错误结果如下:
| colA | colB | Date | ROWNUM |
|---|---|---|---|
| A | 1 | 01-01-2014 | 1 |
| A | 1 | 01-02-2014 | 2 |
| A | 3 | 30-04-2014 | 1 |
| A | 3 | 05-05-2014 | 2 |
| A | 2 | 25-05-2014 | 1 |
| A | 1 | 06-06-2014 | 3 |
| A | 1 | 21-07-2014 | 4 |
| B | 1 | 04-09-2014 | 1 |
| B | 1 | 19-10-2014 | 2 |
| B | 1 | 03-12-2014 | 3 |
| C | 3 | 17-01-2015 | 1 |
| C | 2 | 03-03-2015 | 1 |
| C | 2 | 17-04-2015 | 2 |
解决方案
核心思路是先标记连续相同colA-colB组合的分组,再在每个分组内生成序号,具体实现代码如下:
from pyspark.sql import Window from pyspark.sql.functions import lag, sum, when, row_number # 假设原始DataFrame名为df # 定义排序窗口:按colA分区,Date排序 window_order = Window.partitionBy("colA").orderBy("Date") # 1. 标记分组中断点:当前行与上一行colB不同时标记为1,否则0 df_with_flag = df.withColumn( "group_flag", when( (lag("colB", 1).over(window_order) != df["colB"]) | (lag("colB", 1).over(window_order).isNull()), 1 ).otherwise(0) ) # 2. 生成连续分组ID:累加中断点标记,得到每个连续分组的唯一ID window_group = Window.partitionBy("colA").orderBy("Date").rowsBetween(Window.unboundedPreceding, 0) df_with_group = df_with_flag.withColumn("group_id", sum("group_flag").over(window_group)) # 3. 在每个连续分组内生成序号 window_rownum = Window.partitionBy("colA", "group_id").orderBy("Date") result_df = df_with_group.withColumn("ROWNUM", row_number().over(window_rownum)).drop("group_flag", "group_id") # 查看最终结果 result_df.show()
代码说明
lag("colB",1).over(window_order):获取当前行在colA分区内的上一行colB值,用于判断分组是否中断。group_flag:当分组中断(上一行colB不同或为分区内第一行)时标记为1,否则为0。group_id:对group_flag累加求和,同一个连续组合会得到相同的ID,新组合的ID自动递增。- 最后通过
row_number()基于colA和group_id分区生成序号,实现连续组合的序号重置。
内容的提问来源于stack exchange,提问作者Minal
相关产品推荐
相关产品推荐

