Spark 2.2中如何基于指定列表过滤DataFrame的数组列
解决Spark 2.2中数组元素过滤的问题
嘿,我来帮你搞定这个数组过滤的需求!你之前遇到的问题确实是因为Spark的数组类型和Scala原生List的类型不匹配导致的,我给你两种靠谱的解决方案:
方案一:修正自定义UDF
Spark在DataFrame里存储的数组类型,实际对应的是scala.collection.mutable.WrappedArray,而不是你写UDF时用的Scala原生List,这就是为什么你的UDF会报错。我们只需要调整参数类型就能解决:
import org.apache.spark.sql.functions.udf import scala.collection.mutable.WrappedArray // 你的过滤列表 val filter_list = List("item1", "item2") // 修正后的UDF:参数类型改为WrappedArray[String] val filterItemsUdf = udf { (items: WrappedArray[String]) => // 过滤出在列表中的元素,最后转成Array保持和原字段类型一致 items.filter(filter_list.contains(_)).toArray } // 应用UDF到你的数据 val rawData = Seq(("id1",Array("item1","item2","item3","item4")), ("id2",Array("item1","item2","item3"))) val data = spark.createDataFrame(rawData).toDF("id", "items") val filteredData = data.withColumn("items", filterItemsUdf($"items")) filteredData.show()
方案二:用Spark内置函数实现(无UDF)
如果你不想写UDF,也可以用Spark 2.2支持的内置函数组合来实现,这种方式更符合Spark的分布式优化逻辑:
import org.apache.spark.sql.functions.{explode, collect_list} val rawData = Seq(("id1",Array("item1","item2","item3","item4")), ("id2",Array("item1","item2","item3"))) val data = spark.createDataFrame(rawData).toDF("id", "items") val filter_list = List("item1", "item2") val filteredData = data // 把数组拆分成多行 .select($"id", explode($"items").alias("item")) // 过滤出在目标列表中的元素 .filter($"item".isin(filter_list:_*)) // 按id分组,重新聚合为数组 .groupBy($"id") .agg(collect_list($"item").alias("items")) filteredData.show()
两种方案的输出结果
不管用哪种方法,最终都会得到你想要的结果:
+---+------------+ | id| items| +---+------------+ |id1|[item1, item2]| |id2|[item1, item2]| +---+------------+
小提示:如果你的过滤逻辑比较复杂,UDF更灵活;如果是简单的元素匹配,用内置函数的方式性能会更好,因为Spark能对内置函数做更多优化~
内容的提问来源于stack exchange,提问作者Andreyn
相关产品推荐
相关产品推荐

