如何在Scala-Spark中计算数组类型列的逐元素平均值?
解决Spark数组列逐元素平均值计算问题
要实现数组列的逐元素平均值计算,我们可以分两步走:先将字符串数组转换为数值类型数组,再针对数组的每个位置计算所有行的平均值,最后重新组合成结果数组。以下是具体实现方案:
步骤1:转换数组元素类型
首先需要把示例中的字符串数组转为Double类型数组,这样才能进行数值计算:
import org.apache.spark.sql.functions._ import org.apache.spark.sql.types.DoubleType // 原始数据 val x = Seq("2 4 6", "0 0 0").toDF("value").withColumn("value", split($"value", " ")) // 将字符串数组转换为Double类型数组 val xWithDoubles = x.withColumn("value", transform($"value", elem => elem.cast(DoubleType)))
步骤2:计算逐元素平均值
接下来我们需要对数组的每个索引位置求所有行的平均值,再组合成结果数组。这里分两种情况处理:
情况1:已知数组长度(比如你的实际场景是512)
如果提前知道数组长度,可以直接遍历每个索引生成平均表达式:
// 替换为你的实际数组长度,比如512 val arrayLength = 3 // 生成每个索引位置的平均值表达式 val avgPerIndex = (0 until arrayLength).map(index => avg(col("value")(index))) // 聚合计算并组合成结果数组 val result = xWithDoubles.agg( array(avgPerIndex: _*).alias("avg_value") ) result.show(false)
执行后就能得到你期望的结果:
+---------+ |avg_value| +---------+ |[1.0,2.0,3.0]| +---------+
情况2:未知数组长度
如果数组长度不确定,可以先从数据中动态获取:
// 从第一行数据获取数组长度 val arrayLength = xWithDoubles.select(size($"value")).first().getInt(0) // 后续步骤和情况1一致 val avgPerIndex = (0 until arrayLength).map(index => avg(col("value")(index))) val result = xWithDoubles.agg(array(avgPerIndex: _*).alias("avg_value"))
原理说明
transform函数:遍历数组的每个元素,完成字符串到Double的类型转换。col("value")(index):通过索引访问数组的特定位置元素。avg聚合函数:对所有行的同一索引位置元素求平均值。array函数:将多个单列的平均值结果重新组合成数组类型的列。
内容的提问来源于stack exchange,提问作者wainhuang
相关产品推荐
相关产品推荐

