Spark中stddev_pop()与avg()聚合函数返回类型差异问题咨询
Spark Decimal列计算标准差与平均值的类型问题解答
问题1:为什么stddev_pop()和avg()的输出类型不同?
Spark中这两个聚合函数的类型差异由计算逻辑和精度特性决定:
- avg()函数:针对Decimal类型输入,Spark会保留Decimal的精确计算特性。平均值是累加总和(Decimal类型)除以计数(整数),结果会自然扩展小数位数(如输入为Decimal(38,4),输出变为Decimal(38,8)),全程保证计算精度不丢失。
- stddev_pop()函数:标准差计算涉及平方、求和、开平方等浮点运算。Spark实现中为兼顾通用性与性能,默认返回Double类型——Decimal类型的开平方操作精度处理复杂度高,而Double的精度足以满足绝大多数标准差计算场景需求。
问题2:如何统一两列的格式?
有两种常用方案,可根据实际需求选择:
方案1:将标准差转为Decimal类型(推荐用于数据存储/后续精确处理)
使用cast()函数将stddev_pop结果转为指定Decimal类型(如Decimal(38,4)),与原列精度对齐:
import org.apache.spark.sql.types.DecimalType val df1 = df.groupBy(col("key")) .agg( stddev_pop("count").cast(DecimalType(38,4)).as("std dev"), avg("count").as("average") )
转换后的结果结构:
root |-- key: string (nullable = false) |-- std dev: decimal(38,4) (nullable = true) |-- average: decimal(38,8) (nullable = true)
输出表格示例:
| key | std dev | average |
|---|---|---|
| 2_AN8068571086_EPA_EUR_PID1742804_ik | 3.4993 | 4.57142900 |
若需要让average也统一到4位小数,可进一步对其做类型转换:
avg("count").cast(DecimalType(38,4)).as("average")
方案2:格式化Double类型的显示(仅用于可视化输出)
如果仅需展示时统一小数位数,无需修改数据类型,可使用format_number()函数格式化输出:
import org.apache.spark.sql.functions.format_number val df1 = df.groupBy(col("key")) .agg( format_number(stddev_pop("count"), 4).as("std dev"), format_number(avg("count"), 4).as("average") )
注意:该方式会将结果转为String类型,适合展示,但不适合后续数值计算。
可复现问题的完整代码
import org.apache.spark.sql.Row import org.apache.spark.sql.types._ val schema = StructType( Seq( StructField("key", StringType, nullable = false), StructField("count", DecimalType(38,4), nullable = false) ) ) val data = Seq( Row("2_AN8068571086_EPA_EUR_PID1742804_ik", BigDecimal(2.0)), Row("2_AN8068571086_EPA_EUR_PID1742804_ik", BigDecimal(10.0)), Row("2_AN8068571086_EPA_EUR_PID1742804_ik", BigDecimal(2.0)), Row("2_AN8068571086_EPA_EUR_PID1742804_ik", BigDecimal(4.0)), Row("2_AN8068571086_EPA_EUR_PID1742804_ik", BigDecimal(2.0)), Row("2_AN8068571086_EPA_EUR_PID1742804_ik", BigDecimal(2.0)), Row("2_AN8068571086_EPA_EUR_PID1742804_ik", BigDecimal(10.0)) ) val df = spark.createDataFrame(spark.sparkContext.parallelize(data), schema) df.printSchema() df.show(false) // 类型转换版本示例 val df1 = df.groupBy(col("key")) .agg( stddev_pop("count").cast(DecimalType(38,4)).as("std dev"), avg("count").as("average") ) df1.printSchema() df1.show(false)
内容的提问来源于stack exchange,提问作者M. Yousfi
相关产品推荐
相关产品推荐

