基于PySpark DataFrame按ID及日期生成分类列累积序列
为PySpark DataFrame分类变量生成分组累积序列
问题描述
我们需要基于PySpark DataFrame实现以下需求:
- 按
id字段分组 - 每组内按
date字段排序 - 为每个分类字段(如
cat1、cat2、cat3)生成从该组最早日期记录到当前记录的空格分隔累积序列 - 保留所有原始记录,不能丢失任何行
输入示例
+----+------------------------------+------+-------+------+ | id | date | cat1 | cat2 | cat3 | +----+------------------------------+------+-------+------+ | 1 | 2018-01-25 00:00:... | C | Text1 | val1 | | 1 | 2018-01-25 00:00:... | A | Text1 | val3 | | 1 | 2018-01-25 00:00:... | B | Text5 | val5 | | 2 | 2018-01-26 00:00:... | A | Text2 | val1 | | 2 | 2018-01-26 00:00:... | A | Text1 | val2 | | 3 | 2018-01-27 00:00:... | C | Text6 | val1 | | 3 | 2018-01-29 00:00:... | A | Text2 | val9 | | 3 | 2018-01-29 00:00:... | C | Text6 | val5 | | 3 | 2018-02-05 00:00:... | A | Text1 | val3 | +----+------------------------------+------+-------+------+
输出示例
+----+------------------------------+----------+-------------------------+---------------------+ | id | date | cat1_seq | cat2_seq | cat3_seq | +----+------------------------------+----------+-------------------------+---------------------+ | 1 | 2018-01-25 00:00:... | C | Text1 | val1 | | 1 | 2018-01-25 00:00:... | C A | Text1 Text1 | val1 val3 | | 1 | 2018-01-25 00:00:... | C A B | Text1 Text1 Text5 | val1 val3 val5 | | 2 | 2018-01-26 00:00:... | A | Text2 | val1 | | 2 | 2018-01-26 00:00:... | A A | Text2 Text1 | val1 val2 | | 3 | 2018-01-27 00:00:... | C | Text6 | val1 | | 3 | 2018-01-29 00:00:... | C A | Text6 Text2 | val1 val9 | | 3 | 2018-01-29 00:00:... | C A C | Text6 Text2 Text6 | val1 val9 val5 | | 3 | 2018-02-05 00:00:... | C A C A | Text6 Text2 Text6 Text1 | val1 val9 val5 val3 | +----+------------------------------+----------+-------------------------+---------------------+
解决方案代码
利用PySpark的窗口函数可以轻松实现这个需求,具体代码如下:
import pyspark.sql.functions as f from pyspark.sql import Window # 生成每个分类列的累积序列列 df1 = df.select( "*", *[ f.concat_ws( " ", f.collect_list(c).over(Window.partitionBy("id").orderBy("date")) ).alias(f"{c}_seq") for c in ["cat1", "cat2", "cat3"] # 可替换为你需要处理的分类列列表 ] )
代码逻辑解释
- 窗口定义:
Window.partitionBy("id").orderBy("date")按id分组,组内按date排序,窗口默认覆盖从组内第一条记录到当前记录的范围,实现累积计算。 - 收集累积值:
f.collect_list(c)在窗口范围内收集当前列的所有值,形成一个有序列表。 - 拼接成字符串:
f.concat_ws(" ", ...)将收集到的列表用空格拼接成连续字符串,得到符合要求的序列格式。 - 批量生成列:通过列表推导式批量处理所有分类列,避免重复编写代码,提升效率。
内容的提问来源于stack exchange,提问作者Arijit
相关产品推荐
相关产品推荐

