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

PySpark技术问询:按分组匹配条件列值并新增重复列

在PySpark中按分组填充新列的解决方案

嘿,我来帮你搞定这个PySpark的列填充需求!根据你描述的场景,我们需要给每个col1分组的所有行,填充该分组内col3=1对应的col2值——哪怕分组里有多个col3=1的行,也能灵活处理。下面给你两种常用的实现方法,以及不同多值场景的适配方式:

方法一:使用窗口函数(推荐,无需额外Join)

窗口函数可以直接在原DataFrame上计算,步骤更简洁:

首先导入需要的函数:

from pyspark.sql import Window
from pyspark.sql.functions import first, last, avg, concat_ws, collect_list, when, col

然后根据你的需求选择对应的逻辑:

场景1:取分组内第一个col3=1的col2值

# 定义窗口:按col1分组,可根据需要排序确保col3=1的行优先被选中
window_spec = Window.partitionBy("col1").orderBy("col3")

df = df.withColumn(
    "col4",
    # 先把非col3=1的行设为null,再取分组内第一个非null值
    first(when(col("col3") == 1, col("col2")), ignorenulls=True).over(window_spec)
)

场景2:取分组内最后一个col3=1的col2值

把上面的first换成last即可:

df = df.withColumn(
    "col4",
    last(when(col("col3") == 1, col("col2")), ignorenulls=True).over(window_spec)
)

场景3:取所有col3=1的col2值的平均值

df = df.withColumn(
    "col4",
    avg(when(col("col3") == 1, col("col2"))).over(Window.partitionBy("col1"))
)

场景4:把所有col3=1的col2值拼接成字符串

df = df.withColumn(
    "col4",
    concat_ws(",", collect_list(when(col("col3") == 1, col("col2")))).over(Window.partitionBy("col1"))
)

方法二:先聚合再Join(适合大数据量场景)

如果你的数据集非常大,先聚合出每个分组的目标值再Join,可能会有更好的性能:

步骤1:聚合得到每个分组的目标值

# 这里以取第一个col3=1的col2为例,同样可以替换成last、avg等
grouped_df = df.filter(col("col3") == 1).groupBy("col1").agg(
    first("col2").alias("col4")
)

步骤2:和原DataFrame关联

df = df.join(grouped_df, on="col1", how="left")

示例结果(以场景1为例)

处理后的DataFrame会变成这样:

+-----+-----+-----+-----+
|col1 |col2 |col3 |col4 |
+-----+-----+-----+-----+
| A| 17| 1| 17|
| A| 16| 2| 17|
| A| 18| 2| 17|
| A| 30| 3| 17|
| B| 35| 1| 35|
| B| 34| 2| 35|
| B| 36| 2| 35|
| C| 20| 1| 20|
| C| 30| 1| 20|
| C| 43| 1| 20|
+-----+-----+-----+-----+

如果是场景4(拼接字符串),分组C的col4会变成"20,30,43",非常灵活!

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.22 07:42:13