Spark Scala聚合中如何过滤数组列并统计符合条件元素数量?
解决Spark中过滤数组组合并聚合统计的问题
错误原因
你遇到的问题是混淆了Spark Column API和Scala集合API:Spark的ColumnName对象不能直接调用Scala集合的filter方法,必须使用Spark提供的内置数组操作函数来处理列数据。
解决方案
无需编写UDF,直接使用Spark内置的filter、size和sum函数即可实现需求,步骤如下:
- 使用
arrays_zip关联两个数组列 - 用
filter过滤出满足feature1 > 0且feature2 == 0的zip元素 - 用
size计算每行符合条件的元素数量 - 分组聚合求和得到总数
完整代码示例:
import org.apache.spark.sql.functions.{arrays_zip, filter, size, sum} // 初始化测试数据(修改部分数据用于验证效果) val df = Seq( ("id1", Array(0,1,2), Array(2,0,4)), // 包含一组符合条件的组合 ("id2", Array(0,1,2), Array(2,3,4)), ("id3", Array(0,1,2), Array(2,3,0)) // 包含一组符合条件的组合 ).toDF("id", "feature1", "feature2") // 关联两个数组列 val dfz = df.withColumn("zipped", arrays_zip($"feature1", $"feature2")) // 计算每行符合条件的数量并聚合 val result = dfz .withColumn("valid_count", size(filter($"zipped", elem => elem.getField("feature1") > 0 && elem.getField("feature2") === 0))) .groupBy("id") // 注意:你原代码中的"query"是笔误,原DataFrame中对应的列是"id" .agg(sum($"valid_count").alias("total_valid_matches")) // 查看结果 result.show()
结果说明
运行上述代码后,输出结果如下:
+---+-------------------+ | id|total_valid_matches| +---+-------------------+ |id1| 1| |id2| 0| |id3| 1| +---+-------------------+
关键函数解释
filter(arrayColumn, condition):Spark内置的数组过滤函数,对数组列中的每个元素应用条件判断,返回符合条件的元素组成的新数组size(arrayColumn):返回数组列的元素个数sum(column):对指定列的值进行聚合求和
内容的提问来源于stack exchange,提问作者Callie
相关产品推荐
相关产品推荐

