Spark Scala中遍历DataFrame嵌套数组列并扁平化取值的最佳实现方法
Spark嵌套数组字段全量遍历取值最优方案
优先选择Spark 2.4及以上版本支持的内置高阶函数实现,全程行内处理无shuffle开销,性能远高于explode拆分再聚合的方案。
核心实现逻辑
- 用
transform遍历Species数组的每一个元素,取出每个元素下mammal数组的所有description值,得到二维字符串数组 - 用
flatten将二维数组打平为一维字符串数组 - 用
array_join将一维数组拼接为逗号分隔的字符串
代码示例
import org.apache.spark.sql.functions._ val resultDf = df.select( array_join( // 过滤空的description值,不需要可以去掉这层filter filter( flatten(expr("transform(Animal.Species, s -> transform(s.mammal, m -> m.description))")), desc => desc.isNotNull ), ", " ).alias("all_descriptions") )
低版本Spark兼容方案
如果使用的Spark版本低于2.4,没有高阶函数支持,可以用explode拆分+分组聚合实现,注意需要保留原表的主键字段用于分组:
import org.apache.spark.sql.functions._ // 替换your_primary_key为你原表中用于唯一标识行的主键字段 val resultDf = df.withColumn("species_item", explode(col("Animal.Species"))) .withColumn("mammal_item", explode(col("species_item.mammal"))) .select(col("your_primary_key"), col("mammal_item.description")) .groupBy("your_primary_key") .agg(concat_ws(", ", collect_list(col("description"))).alias("all_descriptions"))
内容的提问来源于stack exchange,提问作者Defcon
相关产品推荐
相关产品推荐

