Scala Spark:如何不用expr改用when条件重构多列过滤代码
解决方案
你可以通过抽象重复逻辑+利用Spark的高阶函数/集合操作来避免冗余代码,下面是两种简洁的实现方式,完全等价于你原来的expr逻辑:
方式一:用foldLeft+条件OR拼接
先把要匹配的类型和分组列定义成列表,然后通过循环构建条件链:
import org.apache.spark.sql.functions.{col, when, lit} import org.apache.spark.sql.types.StringType // 定义需要匹配的类型集合和分组列名 val targetTypes = List("TYPE_A", "TYPE_B", "TYPE_C", "TYPE_D") val groupColumns = List("GROUP_A", "GROUP_B", "GROUP_C", "GROUP_D") // 用foldLeft逐步构建when条件链 val mainTypeExpr = targetTypes.foldLeft(lit("0").cast(StringType)) { (acc, currentType) => // 对每个类型,检查所有分组列是否有等于当前类型的(等价于原expr的in逻辑) val typeCondition = groupColumns.map(col(_) === lit(currentType)).reduce(_ || _) when(typeCondition, lit(currentType)).otherwise(acc) } // 应用到DataFrame val outDF = originalDF.withColumn("MAIN_TYPE", mainTypeExpr)
方式二:用array+array_contains简化条件判断
把所有分组列打包成一个数组,直接用array_contains检查目标类型是否存在,逻辑更直观:
import org.apache.spark.sql.functions.{array, array_contains, when, lit} import org.apache.spark.sql.types.StringType val targetTypes = List("TYPE_A", "TYPE_B", "TYPE_C", "TYPE_D") val groupColumns = List("GROUP_A", "GROUP_B", "GROUP_C", "GROUP_D") // 将分组列转为数组 val groupsArray = array(groupColumns.map(col(_)): _*) // 构建条件链 val mainTypeExpr = targetTypes.foldLeft(lit("0").cast(StringType)) { (acc, currentType) => when(array_contains(groupsArray, lit(currentType)), lit(currentType)).otherwise(acc) } val outDF = originalDF.withColumn("MAIN_TYPE", mainTypeExpr)
为什么这两种方式更优?
- 完全避免了重复写
when语句,不管类型数量m或分组列数量n增加,只需要修改两个列表即可 - 逻辑和原
expr完全一致:按顺序匹配TYPE_A到TYPE_D,优先匹配前面的类型,都不匹配则返回'0' - 代码更易维护、可读性更高
内容的提问来源于stack exchange,提问作者Basil
相关产品推荐
相关产品推荐

