如何用PySpark基于另一DataFrame过滤ArrayType列生成布尔列?
问题
现有两个PySpark DataFrame,结构如下:
df1(Column B为ArrayType(StringType)类型)
| Column A | Column B | Column C | Column D |
|---|---|---|---|
| 1 | [Tokyo, Singapore] | 4 hours | apple |
| 2 | [Tokyo, New York, Paris] | 1.5 hours | banana |
| 3 | [Paris] | 2 hours | orange |
df2
| Destination |
|---|
| Paris |
| New York |
需要在df1中新增一列,规则为:检查Column B数组中的每个元素是否存在于df2的Destination列中,对应位置返回True/False,示例结果如下:
| Column A | Column B | Column C | Column D | new column |
|---|---|---|---|---|
| 1 | [Tokyo, Singapore] | 4 hours | apple | [False, False] |
| 2 | [Tokyo, New York, Paris] | 1.5 hours | banana | [False, True, True] |
| 3 | [Paris] | 2 hours | orange | [True] |
df1的数组长度无上限,df2约有1000行,尝试实现时多次遇到column not iterable类报错,求PySpark实现方法。
解决方案
核心思路
报错原因是直接遍历DataFrame列的错误操作,需利用PySpark高阶函数结合广播变量实现分布式高效处理:
- 将df2的目标值转为Python集合,封装为广播变量(df2数据量小,广播后减少集群数据重复传输)
- 用
transform高阶函数遍历df1的数组列,逐个判断元素是否在广播集合中
代码实现
from pyspark.sql import SparkSession from pyspark.sql.functions import transform, lit # 初始化SparkSession(已初始化可跳过) spark = SparkSession.builder.appName("CheckArrayElements").getOrCreate() # 1. 提取df2的目标值并创建广播变量 dest_values = set(df2.select("Destination").rdd.flatMap(lambda x: x).collect()) broadcast_dest = spark.sparkContext.broadcast(dest_values) # 2. 新增目标列 result_df = df1.withColumn( "new column", transform("Column B", lambda elem: lit(elem.isin(broadcast_dest.value))) ) # 查看结果 result_df.show(truncate=False)
关键说明
- 广播变量:将df2的1000条目标值广播到所有Executor节点,避免重复加载数据,提升处理效率
transform函数:PySpark 3.0+支持的数组高阶函数,专门用于遍历数组元素并应用逻辑,符合分布式计算范式- 禁止直接遍历DataFrame行/列:这类操作违背PySpark分布式设计逻辑,是
column not iterable报错的核心原因
内容的提问来源于stack exchange,提问作者Meg
相关产品推荐
相关产品推荐

