Scala Spark基于多条件过滤DataFrame列的问题求助
嘿,我明白你碰到的麻烦了——直接拿列和数组用===对比肯定跑不通,Spark的Column操作可不认这种写法。不过不用靠循环一个个拼DataFrame,有两个简单直接的方法能一次性搞定!
方法1:用isin()(最省心的方案)
Spark的Column API专门提供了isin()方法,就是用来判断列值是否属于某个集合的,完美匹配你的需求。只需要把数组转成可变参数传入就行:
val stateArray = Array("USAL", "USMD", "USCA", "USME", "USND", "USSD", "USWY", "USAK", "USWA", "USFL", "USGA", "USSC", "USNC", "USMA", "USNH", "USVT", "USAR", "USAZ", "USTX", "USLA", "USIL", "USOR", "USNV", "USID", "USMN", "USNM", "USNE", "USNJ", "USDE", "USVA", "USWV", "USTN", "USKY", "USNY", "USPA", "USIN", "USOH", "USHI", "USOK", "USIA", "USMI", "USMS", "USMO", "USCO", "USKS", "USUT", "USWI", "USMT", "USRI", "USCT") // 一次性过滤出所有符合州代码的行 val tmpDf3 = tmpDf1.filter(tmpDf1("Actor1Geo_ADM1Code").isin(stateArray: _*)) // 验证结果 tmpDf3.show(false) tmpDf3.printSchema()
这里的stateArray: _*是把Scala数组转换成Sparkisin()方法需要的可变参数,这一步不能少哦。
方法2:用array_contains()(适合复杂场景扩展)
如果以后你需要结合其他逻辑做更灵活的过滤,可以用Spark的数组函数来实现:
import org.apache.spark.sql.functions.{array, array_contains, lit} // 把州代码数组转换成Spark的ArrayType列 val statesColumn = array(stateArray.map(lit): _*) // 判断列值是否在这个数组里 val tmpDf3 = tmpDf1.filter(array_contains(statesColumn, tmpDf1("Actor1Geo_ADM1Code")))
这个方法比isin()稍复杂,但扩展性更强,比如可以动态生成数组列或者结合其他条件。
为什么你的原代码会报错?
你原来写的tmpDf("Actor1Geo_ADM1Code") === stateArray,左边是Spark的Column对象,右边是普通的Scala数组——Spark的===操作符只能对比Column和Column,或者Column和单个字面量,没法直接和数组做对比,所以会抛出类型不匹配的错误。
另外提一句:你用for循环的方式虽然能输出每个州的数据,但如果想把这些小DataFrame合并,其实可以用union()(Spark 2.x及以后版本),但这种方式效率远不如直接用isin()——因为isin()是一次性扫描全量数据,而循环union会多次扫描数据源,数据量大的时候性能差距会很明显。
内容的提问来源于stack exchange,提问作者fletchr
相关产品推荐
相关产品推荐

