Spark如何拆分大数组列 实现分组按索引求数组均值
报错原因
代码抛出异常的核心原因是你生成的all_columns是纯字符串列表,不是Spark SQL要求的Column类型表达式。Spark会把select方法中传入的字符串直接解析为列名,自然找不到名为$"weights".getItem(0).as("dim0")的列,触发分析错误。
方案1:动态生成列表达式拆分后聚合
不需要逐行手写300个维度的拆分逻辑,直接在循环中生成Column对象而非代码字符串,即可自动完成数组拆分,后续按常规分组聚合逻辑处理即可:
import org.apache.spark.sql.functions._ val dimTotal = typedConfig.embeddingDims // 动态生成所有维度的拆分列 val splitDimCols = (0 until dimTotal).map(i => col("weights").getItem(i).as(s"dim_$i")) // 生成每个维度的平均值聚合表达式 val aggExprs = splitDimCols.map(c => avg(c).as(s"avg_${c}")) val groupedDf = sample_output .select(col("id") +: splitDimCols: _*) .groupBy("id") .agg(aggExprs.head, aggExprs.tail: _*) // 拼接回数组格式,匹配期望输出结构 val result = groupedDf.select( col("id"), array((0 until dimTotal).map(i => col(s"avg_dim_$i")): _*).as("weights") ) result.show(false)
方案2:高阶函数直接聚合数组(无宽表,性能更好)
Spark 2.4及以上版本支持数组高阶函数,不需要拆分为数百个单独列,直接在数组结构上按索引计算平均值即可,大幅减少宽表带来的序列化、调度开销:
import org.apache.spark.sql.functions._ val result = sample_output .groupBy("id") .agg(collect_list("weights").as("all_weights")) .select( col("id"), expr(""" transform( sequence(0, size(all_weights[0]) - 1), idx -> aggregate( transform(all_weights, arr -> arr[idx]), 0D, (acc, curr) -> acc + curr, total -> total / size(all_weights) ) ) as weights """) ) result.show(false)
如果使用Spark 3.0及以上版本,内置了array_avg函数,可以简化expr中的逻辑:
transform( sequence(0, size(all_weights[0]) - 1), idx -> array_avg(transform(all_weights, arr -> arr[idx])) ) as weights
注意:两种方案都要求同组内所有weights数组长度一致,否则会出现索引越界或计算结果为空的问题,存在脏数据时建议提前过滤长度不符合要求的记录。
内容的提问来源于stack exchange,提问作者Ethan Seiler
相关产品推荐
相关产品推荐

