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
相关产品推荐
相关产品推荐

