Spark Scala:DataFrame新增列统计指定列中1的个数问题求助
问题原因分析
你遇到的问题核心是误用了Scala原生的if语句处理Spark Column对象——Spark的Column是分布式计算的列表达式,原生if无法识别这种列级别的条件判断,它只会在Driver端执行一次静态判断,不会对每行数据做逐行处理,所以最终cnt始终是初始的0,导致ones列全为0。
解决方案
我们需要用Spark提供的列级条件函数when来实现逐行判断,下面是修正后的完整代码:
import org.apache.spark.sql.functions._ import sqlContext.implicits._ import org.apache.spark.sql.Column val input = sc.parallelize(Seq( ("r1", 1, 1), ("r2", 6, 4), ("r3", 4, 1), ("r4", 1, 2) )).toDF("ID", "a", "b") def countOnes(cols: Column*) = cols.foldLeft(lit(0)){ (cnt, current) => // 用Spark的when函数做列级条件判断,满足则加1,否则加0 cnt + when(current === 1, 1).otherwise(0) } val output = input.withColumn("ones", countOnes(col("a"), col("b"))) output.show()
运行这段代码后,就能得到你预期的结果:
+---+---+---+----+ | ID| a| b|ones| +---+---+---+----+ | r1| 1| 1| 2| | r2| 6| 4| 0| | r3| 4| 1| 1| | r4| 1| 2| 1| +---+---+---+----+
另一种简洁实现思路
你也可以把每个列的判断结果转成整数(布尔值true转1,false转0),然后直接求和,代码更紧凑:
import org.apache.spark.sql.types.IntegerType def countOnes(cols: Column*) = cols .map(c => (c === 1).cast(IntegerType)) // 把判断结果转成整数 .reduce(_ + _) // 对所有列的结果求和
这个逻辑和上面的实现效果完全一致,选哪种都可以~
内容的提问来源于stack exchange,提问作者loba76
相关产品推荐
相关产品推荐

