You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.27 20:06:28