Spark DataFrame聚合时如何对收集的数组元素进行过滤
需求说明
对Spark DataFrame执行聚合操作,获取每个广告主对应的品牌数组,同时参照给定的品牌白名单表过滤数组,仅保留白名单中存在的品牌。
原始广告主品牌表(df2)
+------------+------+ |advertiser |brand | +------------+------+ |Advertiser 1|Brand1| |Advertiser 1|Brand2| |Advertiser 2|Brand3| |Advertiser 2|Brand4| |Advertiser 3|Brand5| |Advertiser 3|Brand6| +------------+------+
原有聚合逻辑
import org.apache.spark.sql.functions.collect_list df2 .groupBy("advertiser") .agg(collect_list("brand").as("brands"))
原有逻辑输出结果:
+------------+----------------+ |advertiser |brands | +------------+----------------+ |Advertiser 1|[Brand1, Brand2]| |Advertiser 2|[Brand3, Brand4]| |Advertiser 3|[Brand5, Brand6]| +------------+----------------+
过滤用品牌白名单表
+------+------------+ |brand |brand name | +------+------------+ |Brand1|Brand_name_1| |Brand3|Brand_name_3| +------+------------+
期望输出结果
+------------+--------+ |advertiser |brands | +------------+--------+ |Advertiser 1|[Brand1]| |Advertiser 2|[Brand3]| |Advertiser 3|null | +------------+--------+
实现方案
提供两种可选择的实现方式,可根据实际场景选择:
方案1:先过滤再聚合(推荐,大数据量下性能更优)
先通过左连接标记有效品牌,再聚合仅收集有效品牌,避免无效数据参与shuffle,性能更好。
假设白名单表变量名为brand_whitelist_df,代码如下:
import org.apache.spark.sql.functions.{col, collect_list, when, size} val result = df2 // 左连接白名单表,标记当前品牌是否属于白名单 .join( brand_whitelist_df.select("brand").withColumnRenamed("brand", "valid_brand"), df2("brand") === col("valid_brand"), "left" ) .groupBy("advertiser") // 仅收集白名单内的品牌 .agg(collect_list(when(col("valid_brand").isNotNull, col("brand"))).as("brands")) // 空数组替换为null,匹配预期输出格式 .withColumn("brands", when(size(col("brands")) === 0, null).otherwise(col("brands")))
方案2:聚合后过滤(适合已完成聚合的场景)
如果已经得到聚合后的结果,不想重新执行全量聚合流程,可以通过数组交集函数实现过滤:
import org.apache.spark.sql.functions.{collect_list, array_intersect, lit, when, size} import scala.collection.JavaConverters._ // 收集白名单品牌为数组,广播到所有计算节点 val validBrands = spark.sparkContext.broadcast( brand_whitelist_df.select("brand").as[String].collect().toSeq.asJava ) val result = df2 .groupBy("advertiser") .agg(collect_list("brand").as("brands")) // 取品牌数组和白名单的交集 .withColumn("brands", array_intersect(col("brands"), lit(validBrands.value))) // 空数组替换为null .withColumn("brands", when(size(col("brands")) === 0, null).otherwise(col("brands")))
注意:如果不需要将空数组转为null,可以去掉最后一步
withColumn的转换逻辑。
内容的提问来源于stack exchange,提问作者Meesam
相关产品推荐
相关产品推荐

