PySpark嵌套循环过滤DataFrame的并行化实现及报错解决问询
我看过许多同类问题,但仍有困惑:有人建议用Groupby,有人提议使用map或flatmap,不确定该尝试哪种方案。
输入为包含date、term、brand、text列的PySpark DataFrame,期望输出为过滤后的DataFrame列表[df1, df2, ...]。原有Python嵌套循环代码如下(df为约10000条记录的大型DataFrame):
filtered_df_list = [] for term in term_list: for br in brand_list: for dt in date_list: tcd_df = df.filter( (df.term == term) & (df.brand == brand) & (df.date == dt) ) if len(tcd_df.index) > 0: filtered_df_list.append(tcd_df)
另外,能否基于term、brand、date列表创建DataFrame,并通过withColumn新增存储过滤后DataFrame的列?同时需要并行化实现的入门代码。
尝试过的错误代码及报错
我尝试了以下代码,但工作节点尝试用SparkContext执行过滤时出现不允许的操作:
rdd2 = df.rdd.map(lambda row: ( row["date"], row["brand"], row["term"], row["text"], df.filter((df.brand == row["brand"]) & (df.term == row["term"]) & (df.date == row["date"])) )) print("\n\nAfter loop") df2 = rdd2.toDF(["date", "brand", "term", "text"]) df2.show()
报错信息翻译:
_pickle.PicklingError: 无法序列化对象:RuntimeError:看起来你正尝试从广播变量、Action或Transformation中引用SparkContext。SparkContext只能在Driver端使用,不能在Worker端运行的代码中使用。更多信息请参考SPARK-5063。
1. 优化原有嵌套循环(简单直接)
原有循环的问题是每次过滤都会触发一次Spark作业,效率低。可以先缩小数据集范围,再拆分分组:
from itertools import product # 生成所有(term, brand, date)的组合 filter_combinations = list(product(term_list, brand_list, date_list)) # 先一次性过滤出所有符合条件的数据,减少后续扫描次数 filtered_df = df.filter( (df.term.isin(term_list)) & (df.brand.isin(brand_list)) & (df.date.isin(date_list)) ) filtered_df_list = [] for term, br, dt in filter_combinations: group_df = filtered_df.filter( (filtered_df.term == term) & (filtered_df.brand == br) & (filtered_df.date == dt) ) # 用count()替代len(index),避免全量加载数据 if group_df.count() > 0: filtered_df_list.append(group_df)
2. 用GroupBy实现高效分组(推荐)
Spark的GroupBy是分布式执行的,更适合大数据场景:
from pyspark.sql import functions as F # 按term, brand, date分组,收集每组的所有行 grouped = filtered_df.groupBy("term", "brand", "date").agg( F.collect_list(F.struct("text")).alias("rows") ) # 将每组转换成独立的DataFrame filtered_df_list = [] for row in grouped.collect(): group_data = [(row.term, row.brand, row.date, r.text) for r in row.rows] group_df = spark.createDataFrame(group_data, schema=df.schema) filtered_df_list.append(group_df)
这种方式只需要扫描原DF两次,比循环多次过滤高效得多。
3. 创建包含分组数据列的DataFrame
Spark列无法直接存储另一个DataFrame(无法序列化到Worker端),但可以存储分组的结构化数据,后续再转换为DF:
# 创建包含所有(term, brand, date)组合的DF combination_df = spark.createDataFrame(filter_combinations, schema=["term", "brand", "date"]) # 内连接过滤后的DF,获取每组数据 joined_df = combination_df.join(filtered_df, on=["term", "brand", "date"], how="inner") # 分组收集每组数据到数组列 result_df = joined_df.groupBy("term", "brand", "date").agg( F.collect_list(F.struct("text")).alias("group_data") ) # 在Driver端将group_data转换为DF列(仅Driver端可用) def convert_to_df(row): return (row.term, row.brand, row.date, spark.createDataFrame(row.group_data, schema=["text"])) result_with_df = spark.createDataFrame( result_df.rdd.map(convert_to_df), schema=["term", "brand", "date", "filtered_df"] )
注意:filtered_df列仅能在Driver端操作,Worker端无法使用。
4. 并行化实现入门代码
Spark本身已经是分布式并行执行的,若要在Driver端并行处理分组结果,可使用Python多进程:
from multiprocessing import Pool def process_group(group_row): term, br, dt, rows = group_row # 这里仅做本地数据处理,不调用Spark API return (term, br, dt, len(rows)) # 把分组数据拉到Driver端 grouped_data = [(row.term, row.brand, row.date, row.rows) for row in grouped.collect()] # 并行处理分组 with Pool(4) as p: results = p.map(process_group, grouped_data) # 转换为Spark DF result_df = spark.createDataFrame(results, schema=["term", "brand", "date", "row_count"])
错误原因说明
你尝试的代码报错,是因为在rdd.map(Worker端执行的Transformation)中引用了原DataFramedf,而df依赖SparkContext。SparkContext无法序列化到Worker端,所以不能在Transformation里调用需要SparkContext的API(比如filter、createDataFrame)。
内容的提问来源于stack exchange,提问作者user1717931

