You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PySpark按ID分组聚合连续pos对应的pos与value数组

Spark按ID分组并聚合连续pos的数组

原始DataFrame

+---+-----+------------+
| ID|  pos|       value|
+---+-----+------------+
| A1|    0|         ABC|
| A1|    1|         BCD|
| A1|    2|         CDF|
| A1|    5|         ABC|
| A1|    8|         FGR|
| A1|    9|         EFD|
| A2|    0|         L_1|
| A2|    1|         L_2|
| A2|    3|         STU|
+---+-----+------------+

期望输出

+---+-------------+-----------------------------+
| ID|     arr(pos)|                   arr(value)|
+---+-------------+-----------------------------+
| A1|    [0, 1, 2]|        ['ABC', 'BCD', 'CDF']|
| A1|          [5]|                      ['ABC']|
| A1|       [8, 9]|               ['FGR', 'EFD']|
| A2|       [0, 1]|               ['L_1', 'L_2']|
| A2|          [3]|                      ['STU']|
+---+-------------+-----------------------------+

核心需求

  • 按ID分组,对pos和value分别聚合为数组
  • 同一ID下,pos数组中的数值必须是连续整数;若pos与前后数值不连续(如A1的pos=5),则单独成组
  • 同一ID的pos值无重复

已尝试的方法

方法1:使用Lag函数识别前后pos

代码:

from pyspark.sql import functions as F, Window as W
window = W.partitionBy("ID").orderBy("pos")

df = (
    df
    .withColumn("prev_pos", F.lag(F.col('pos'), 1).over(window))
    .withColumn("next_pos", F.lead(F.col('pos'), 1).over(window))  # 注:原代码写错,lag(-1)应为lead
)

执行结果:

+---+-----+------------+---------+---------+
| ID|  pos|       value| prev_pos| next_pos|
+---+-----+------------+---------+---------+
| A1|    0|         ABC|     null|        1|
| A1|    1|         BCD|        0|        2|
| A1|    2|         CDF|        1|        5|
| A1|    5|         ABC|        2|        8|
| A1|    8|         FGR|        5|        9|
| A1|    9|         EFD|        8|     null|
| A2|    0|         L_1|     null|        1|
| A2|    1|         L_2|        0|        3|
| A2|    3|         STU|        1|     null|
+---+-----+------------+---------+---------+

问题:无法进一步识别连续区间并分组聚合

方法2:直接分组聚合所有pos

代码:

df = (
    df
    .groupby("ID")
    .agg(
        F.sort_array(F.collect_list(F.struct("pos", "value")))
        .alias("pos_list")
    )
    .withColumn("arr(pos)", F.col("pos_list.pos"))
    .withColumn("arr(value)", F.col("pos_list.value"))
    .drop("pos_list")
)

执行结果:

+---+--------------------+---------------------------------------------+
| ID|            arr(pos)|                                   arr(value)|
+---+--------------------+---------------------------------------------+
| A1|  [0, 1, 2, 5, 8, 9]|   ['ABC', 'BCD', 'CDF', 'ABC', 'FGR', 'EFD']|
| A2|           [0, 1, 3]|                        ['L_1', 'L_2', 'STU']|
+---+--------------------+---------------------------------------------+

问题:无法在pos不连续处拆分数组


解决方案

核心思路:先为每个连续的pos区间标记分组ID,再按ID和区间分组ID聚合。

完整代码:

from pyspark.sql import functions as F, Window as W

# 1. 定义窗口:按ID分区,pos排序
window = W.partitionBy("ID").orderBy("pos")

# 2. 标记连续区间的起点:第一个元素或与前一个pos差值≠1的位置
df_with_gap = df.withColumn(
    "is_new_group",
    F.when(
        F.lag("pos").over(window).isNull(),
        1
    ).when(
        F.col("pos") - F.lag("pos").over(window) != 1,
        1
    ).otherwise(0)
)

# 3. 累积求和生成区间分组ID
df_with_group = df_with_gap.withColumn(
    "group_id",
    F.sum("is_new_group").over(window.rangeBetween(W.unboundedPreceding, 0))
)

# 4. 按ID和group_id聚合,收集数组
result = df_with_group.groupBy("ID", "group_id").agg(
    F.collect_list("pos").alias("arr(pos)"),
    F.collect_list("value").alias("arr(value)")
).drop("group_id").orderBy("ID", F.min("pos").over(W.partitionBy("ID")))

# 展示结果
result.show(truncate=False)

执行结果与期望输出一致:

+---+---------+----------------+
|ID |arr(pos) |arr(value)      |
+---+---------+----------------+
|A1 |[0,1,2]  |[ABC, BCD, CDF] |
|A1 |[5]      |[ABC]           |
|A1 |[8,9]    |[FGR, EFD]      |
|A2 |[0,1]    |[L_1, L_2]      |
|A2 |[3]      |[STU]           |
+---+---------+----------------+

内容的提问来源于stack exchange,提问作者huy

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.21 18:14:55