Scala Spark中expr+filter非UDF过滤时传入预定义列表的方法
问题说明
需要过滤Spark DataFrame的数组类型列,剔除指定排除列表badList中的元素,要求不使用UDF,且不硬编码排除值到SQL表达式中,运行环境为Scala 2.11。
报错原因
直接将Scala原生List传入expr字符串插值的写法会报错,是因为List默认toString方法会生成List(xxx, yyy)格式的内容,完全不符合SQL中IN子句要求的('v1','v2')语法格式,SQL解析器无法识别该片段。
可行实现方案
方案1:拼接合法SQL子句(兼容Spark 2.4+,与原UDF行为完全一致)
先将排除列表处理为SQL语法要求的单引号包裹、逗号分隔的格式,再传入filter高阶函数表达式。该方案不会修改原数组的重复元素、也不会丢弃null值,和你之前写的UDF逻辑完全匹配:
import org.apache.spark.sql.functions.{col, expr} // 处理特殊字符:将值内的单引号转义为两个单引号,符合SQL语法规范 val inClause = badList .map(v => s"'${v.replace("'", "''")}'") .mkString(",") // 空列表边界处理,避免生成IN()非法语法 val posColumn = if (badList.isEmpty) col("allVals") else expr(s"filter(allVals, val -> val NOT IN ($inClause))") val cleanedDF = myDF.withColumn("pos", posColumn)
如果需要同时过滤数组中的null值,只需在lambda条件中补充val is not null判断即可。
方案2:使用array_except内置函数(代码最简洁,结果会去重)
如果业务场景不需要保留数组内的重复元素,也不需要保留null值,可以直接用内置数组差集函数,不需要编写lambda逻辑:
import org.apache.spark.sql.functions.{array_except, array, lit, col} // 将Scala本地列表转为Spark支持的数组字面量 val badArrayLit = array(badList.map(lit(_)): _*) val cleanedDF = myDF.withColumn("pos", array_except(col("allVals"), badArrayLit))
注意:array_except会对返回结果自动去重,且会过滤原数组中的null值
方案3:使用Dataset高阶函数API(仅Spark 3.0+支持,类型安全)
如果后续升级到Spark 3.0及以上版本,可以直接使用DataFrame原生API的数组过滤方法,不需要拼接SQL字符串,类型校验更严格:
import org.apache.spark.sql.functions.col // 排除列表元素较多时建议广播,提升执行效率 val badSet = spark.sparkContext.broadcast(badList.toSet) val cleanedDF = myDF.withColumn( "pos", org.apache.spark.sql.functions.filter(col("allVals"), v => !badSet.value.contains(v)) )
内容的提问来源于stack exchange,提问作者Thal
相关产品推荐
相关产品推荐

