如何高效实现PySpark分组数据的行迭代与配对逻辑
PySpark实现配对出入记录高效处理
数据格式
| id | location | type | date | time | |----|----------|------|-----------|-------| | 1 | 33 | out | 2020-11-03| 08:35 | | 1 | 34 | in | 2020-11-03| 08:37 | | 1 | 33 | in | 2020-11-03| 09:40 | | 1 | 33 | out | 2020-11-03| 10:35 | | 1 | 33 | in | 2020-11-03| 12:40 | | 1 | 33 | out | 2020-11-03| 18:35 | | 2 | 33 | out | 2020-11-03| 11:35 | | 2 | 33 | in | 2020-11-03| 18:35 | | 2 | 33 | both | 2020-11-03| 20:35 |
业务逻辑
- 按
id、location、date对记录分组 - 遍历分组内记录,先找第一条
in或both作为入站记录 - 接着找下一条
out或both作为出站记录 both类型规则:若前面无未配对的in则视为in,有未配对的in则视为out- 每组配对的入站、出站记录拆分为单独一行,同一分组有多组配对则生成多行
当前低效代码
你通过collect()将数据拉到Driver端循环处理,大数据量下会触发Driver内存瓶颈,且完全浪费Spark分布式计算能力,这是耗时极长的核心原因:
df = # 你的样本数据 _intype = ('in', 'both') _outtype = ('out', 'both') lst = df.select("id", "location", "date").distinct().collect() in_lst = [] out_lst = [] lst_agg = [] pre_type = None for record in lst: id_val = record.id location_val = record.location date_val = record.date records = df.filter((df.id == id_val) & (df.location == location_val) & (df.date == date_val)).collect() for row in records: if pre_type is None or pre_type == 'OUT': if row.type in _intype: time_in = row.time time_out = None pre_type = 'IN' elif pre_type == 'IN': if row.type in _outtype: time_out = row.time pre_type = 'OUT' if time_out is not None: lst_agg.append((id_val, location_val, date_val, time_in, time_out)) time_in = None time_out = None df_agg_cols = ["id","location","date","time_in","time_out"] df_agg = spark.createDataFrame(data=lst_agg, schema = df_agg_cols)
预期输出
| id | location | date | time_in | time_out | |----|----------|-----------|---------|----------| | 1 | 33 | 2020-11-03| 09:40 | 10:35 | | 1 | 33 | 2020-11-03| 12:40 | 18:35 | | 2 | 33 | 2020-11-03| 18:35 | 20:35 |
高效分布式实现方案
利用Spark窗口函数和状态标记,全程在Executor端分布式处理,避免拉取数据到Driver:
代码实现
from pyspark.sql import Window from pyspark.sql.functions import col, when, spark_sum # 1. 按分组键+时间排序,确保记录按时间顺序处理 df_sorted = df.orderBy("id", "location", "date", "time") # 2. 定义分组窗口 group_window = Window.partitionBy("id", "location", "date").orderBy("time") # 3. 计算分组内未配对的in数量,标记每条记录的实际角色 df_with_state = df_sorted.withColumn( "is_in_candidate", when(col("type").isin("in", "both"), 1).otherwise(0) ).withColumn( "is_out_candidate", when(col("type").isin("out", "both"), 1).otherwise(0) ).withColumn( # 计算当前记录前的未配对in数量:累计in候选数 - 累计out候选数 "unpaired_in_count", spark_sum(col("is_in_candidate")).over(group_window.rowsBetween(Window.unboundedPreceding, -1)) - spark_sum(col("is_out_candidate")).over(group_window.rowsBetween(Window.unboundedPreceding, -1)) ).withColumn( # 标记当前记录的实际角色 "role", when( (col("unpaired_in_count") == 0) & (col("is_in_candidate") == 1), "IN" # 无未配对in,作为入站 ).when( (col("unpaired_in_count") > 0) & (col("is_out_candidate") == 1), "OUT" # 有未配对in,作为出站 ).otherwise( "INVALID" # 无效记录,跳过 ) ) # 4. 筛选有效记录,生成配对组ID df_valid = df_with_state.filter(col("role").isin("IN", "OUT")) pair_window = Window.partitionBy("id", "location", "date").orderBy("time") df_with_pair_id = df_valid.withColumn( "pair_id", spark_sum(when(col("role") == "IN", 1).otherwise(0)).over(pair_window) ) # 5. 透视合并配对记录,生成最终结果 final_df = df_with_pair_id.groupBy("id", "location", "date", "pair_id")\ .pivot("role")\ .agg({"time": "first"})\ .filter(col("IN").isNotNull() & col("OUT").isNotNull())\ .select("id", "location", "date", col("IN").alias("time_in"), col("OUT").alias("time_out"))\ .orderBy("id", "location", "date", "time_in") final_df.show()
逻辑说明
- 排序与窗口:保证每个分组内的记录按时间顺序处理,窗口函数用于计算分组内的累计状态
- 状态标记:通过
unpaired_in_count判断both记录的角色,符合业务规则 - 配对分组:用
pair_id绑定连续的IN-OUT配对,最后通过透视表将配对时间合并为一行 - 分布式处理:所有计算在Executor端完成,彻底避免
collect()带来的性能瓶颈
内容的提问来源于stack exchange,提问作者Ellie
相关产品推荐
相关产品推荐

