如何基于条件拼接列值,创建PySpark DataFrame新列?
解决PySpark生成Tested_devices列的问题
首先,咱们先拆解你遇到的错误:TypeError: when() takes exactly 2 arguments (1 given),这是因为PySpark的when函数必须同时传入判断条件和满足条件时的返回值,你只写了条件df.Tested == 'Y',没给对应的结果,所以触发了报错。
接下来咱们一步步实现你的需求:针对每个Dev_No,把该分组内所有Tested='Y'的model用逗号拼接,填充到每一行的Tested_devices列中。
方案一:窗口函数法(无需关联表,更简洁)
这种方法直接在原DataFrame上通过窗口函数完成聚合,不用额外做表关联,代码更直观:
# 导入需要的工具类和函数 from pyspark.sql import Window from pyspark.sql.functions import col, when, collect_list, concat_ws # 定义窗口:按Dev_No分组,聚合范围是每个分组内的所有行 window_spec = Window.partitionBy("Dev_No") # 添加Tested_devices列 df1 = df.withColumn( # 收集当前Dev_No下所有Tested='Y'的model,用逗号分隔拼接 "Tested_devices", concat_ws(", ", collect_list(when(col("Tested") == "Y", col("model"))).over(window_spec)) ).withColumn( # 把空字符串转为null,和你的目标输出格式一致 "Tested_devices", when(col("Tested_devices") == "", None).otherwise(col("Tested_devices")) ) # 查看最终结果 df1.show(truncate=False)
代码解释
- 窗口定义:
Window.partitionBy("Dev_No")指定了聚合的范围是每个Dev_No对应的所有行。 - collect_list + when:
collect_list(when(col("Tested") == "Y", col("model")))会收集当前分组内所有符合Tested='Y'的model值,不符合条件的会被转为null,而collect_list会自动忽略这些null值,只保留有效数据。 - concat_ws:把收集到的
model列表用,拼接成字符串,最终生成Tested_devices列。
方案二:分组聚合+左连接(适合习惯SQL关联逻辑的场景)
如果你更熟悉分组聚合后关联表的思路,也可以用这种方式:
from pyspark.sql.functions import concat_ws, collect_list # 第一步:分组聚合,得到每个Dev_No对应的Tested_devices agg_df = df.filter(col("Tested") == "Y") \ .groupBy("Dev_No") \ .agg(concat_ws(", ", collect_list("model")).alias("Tested_devices")) # 第二步:左连接回原DataFrame,没有Tested='Y'记录的Dev_No对应列会显示null df1 = df.join(agg_df, on="Dev_No", how="left") df1.show(truncate=False)
这个方案的输出结果和窗口函数法完全一致,你可以根据自己的习惯选择。
为什么你的原代码不对?
你写的a = df.select("Dev_No", "model"), when(df.Tested == 'Y')存在两个问题:
- when函数用法错误:
when必须传入两个参数(判断条件+满足条件的返回值),正确写法是when(col("Tested") == "Y", col("model"))。 - 语法错误:
df.select(...)后面加逗号和when是无效的Python语法,不能拆分写在逗号两侧。
内容的提问来源于stack exchange,提问作者User12345
相关产品推荐
相关产品推荐

