PySpark DataFrame两次过滤触发KeyError,是否为延迟执行导致?
你碰到的这个KeyError确实和Spark的延迟执行以及查询优化有关,而且这属于Spark处理UDF时的预期行为,下面我来详细解释原因和解决办法:
为什么会出现这个问题?
Spark的Catalyst优化器会尝试对查询计划做各种性能优化,比如合并过滤操作、调整执行顺序,但UDF对Spark来说是个"黑盒"——它无法解析UDF内部的逻辑,所以没办法判断第一个过滤UDF会移除哪些行,也无法保证两个过滤操作的执行顺序。
具体到你的场景:虽然你先调用df.where(udf_indict)得到了df1,但Spark的延迟执行特性意味着df1并没有立即计算。当你执行df1.where(udf_bigenough).show()时,Spark会把两个过滤操作合并到同一个查询计划中,可能会对原始数据集的所有行同时应用两个UDF——包括已经被第一个UDF标记为要移除的'c'行。这时候第二个UDF尝试访问mydict_bc.value['c'],自然就触发了KeyError。
哪怕df1.show()已经展示了没有'c'的结果,那也只是Spark临时计算了df1的结果并展示,但并没有把df1的结果持久化下来,后续的查询还是会基于原始数据集重新计算。
解决办法
针对这个问题,有几种可行的解决方案,按推荐程度排序:
1. 用Spark内置函数替代UDF(最优方案)
Spark的内置函数是Catalyst能理解的,优化器可以正确判断过滤逻辑、保证执行顺序,而且性能比UDF好很多。
from pyspark.sql import functions as func from pyspark.sql.types import StringType, SparkContext, SparkSession sc = SparkContext() ss = SparkSession(sc) mydict = {"a": 4, "b": 6} mydict_bc = sc.broadcast(mydict) # 用内置isin做第一个过滤 df = ss.createDataFrame(["a", "b", "c"], StringType()).toDF("name") df1 = df.where(func.col("name").isin(mydict_bc.value.keys())) df1.show() # 用广播字典的lookup结合内置比较做第二个过滤 df2 = df1.where(func.lit(mydict_bc.value)[func.col("name")] > 5) df2.show()
2. 持久化第一个过滤后的结果
如果必须使用UDF,可以通过cache()或persist()将df1的结果持久化到内存或磁盘,这样后续的过滤操作就只会基于已经过滤后的数据集执行:
from pyspark.sql import functions as func from pyspark.sql.types import StringType, BooleanType, SparkContext, SparkSession sc = SparkContext() ss = SparkSession(sc) mydict = {"a": 4, "b": 6} mydict_bc = sc.broadcast(mydict) udf_indict = func.udf(lambda x: x in mydict_bc.value, BooleanType()) udf_bigenough = func.udf(lambda x: mydict_bc.value[x] > 5, BooleanType()) df = ss.createDataFrame(["a", "b", "c"], StringType()).toDF("name") df1 = df.where(udf_indict('name')).cache() # 持久化df1 df1.show() df1.where(udf_bigenough('name')).show() # 现在只会处理df1中的行,不会碰到'c'
注意:如果数据集很大,缓存会占用集群内存,需要根据实际情况调整存储级别(比如persist(StorageLevel.DISK_ONLY))。
3. 合并两个过滤逻辑到同一个UDF
把两个判断逻辑合并到一个UDF里,先检查键是否存在,再判断阈值,避免出现KeyError:
from pyspark.sql import functions as func from pyspark.sql.types import StringType, BooleanType, SparkContext, SparkSession sc = SparkContext() ss = SparkSession(sc) mydict = {"a": 4, "b": 6} mydict_bc = sc.broadcast(mydict) # 合并两个过滤逻辑 udf_combined = func.udf(lambda x: x in mydict_bc.value and mydict_bc.value[x] > 5, BooleanType()) df = ss.createDataFrame(["a", "b", "c"], StringType()).toDF("name") df.filter(udf_combined('name')).show()
这个方案简单直接,但还是依赖UDF,性能不如内置函数方案。
总结
你碰到的KeyError是Spark优化器处理UDF时的预期行为——因为UDF是黑盒,优化器无法保证过滤顺序。优先使用内置函数可以避免这类问题,同时提升查询性能;如果必须用UDF,持久化中间结果或合并逻辑也是有效的解决办法。
内容的提问来源于stack exchange,提问作者Go Erlangen

