Spark 2.4 Scala中过滤数组列元素并重组的优化方案问询
解决Spark 2.4中DataFrame数组列元素过滤的问题
嘿,我来帮你搞定这个问题!你现在的思路方向是对的,但还差最后一步把数据重新聚合回数组;另外还有更高效的方案,不用拆分行就能直接过滤数组元素,我给你详细说说:
方案一:完善你的explode+filter+重组数组思路
你已经完成了拆分行和过滤的步骤,接下来只需要按id分组,用collect_list把过滤后的元素重新拼成数组就行。如果需要严格保持原始数组中的元素顺序,可以结合monotonically_increasing_id记录顺序(你的示例场景里顺序不影响结果的话,直接用collect_list就够了)。代码如下:
import org.apache.spark.sql.functions.{explode, length, collect_list} val res = Seq(("1", Array("x", "xxx", "xx")), ("2", Array("yy", "y", "yyy"))).toDF("id", "letter") val result = res .withColumn("tmp", explode(col("letter"))) .filter(length(col("tmp")) < 3) .groupBy("id") .agg(collect_list("tmp").alias("letter")) result.show()
运行后就能得到你想要的结果:
+---+-------+ | id| letter| +---+-------+ | 1|[x, xx]| | 2|[yy, y]| +---+-------+
方案二:使用Spark 2.4的数组高阶函数(更优方案)
Spark 2.4开始支持数组高阶函数,可以直接对数组列进行过滤操作,不需要拆分行和分组聚合,避免了shuffle操作,性能更高效。你可以用filter高阶函数直接筛选数组中长度小于3的元素,代码更简洁:
方式1:用SQL表达式风格
import org.apache.spark.sql.functions.expr val res = Seq(("1", Array("x", "xxx", "xx")), ("2", Array("yy", "y", "yyy"))).toDF("id", "letter") val result = res .withColumn("letter", expr("filter(letter, elem -> length(elem) < 3)")) result.show()
方式2:用Scala函数式风格
import org.apache.spark.sql.functions.{filter, length} val result = res .withColumn("letter", filter(col("letter"), elem => length(elem) < 3))
这两种写法都能直接得到目标结果,而且性能更好——不需要把数组拆成多行再聚合,适合处理大数据量的场景。
两种方案对比
- 方案一:适合你已经熟悉explode操作的场景,但需要分组shuffle,数据量大时性能会受影响。
- 方案二:直接操作数组列,无shuffle,代码简洁,性能更优,是Spark 2.4及以上版本的推荐方案。
内容的提问来源于stack exchange,提问作者ticapix
相关产品推荐
相关产品推荐

