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

Spark Scala中Python .transform()的替代实现:添加分组均值特征

How to Add Grouped Mean Feature to Spark DataFrame in Scala (Equivalent to Python's transform())

Got it, let's tackle this problem! In Spark Scala, there isn't a direct 1:1 equivalent to Python's transform() method for this specific use case, but we have two straightforward ways to get the same result—adding a class_mean column that shows the average val for each class group across all rows.

This approach is the most concise and efficient, as it avoids extra join operations by computing the grouped mean directly over a window partitioned by your class column.

First, import the necessary Spark functions and window utilities:

import org.apache.spark.sql.expressions.Window
import org.apache.spark.sql.functions.mean

Define your window to partition by the class column:

val classWindow = Window.partitionBy("class")

Then, use withColumn to add the class_mean column by applying the mean function over the window:

val df = Seq( ("a", 1), ("a", 3), ("b", 3), ("b", 7) ).toDF("class", "val")
val resultDf = df.withColumn("class_mean", mean($"val").over(classWindow))

If you run resultDf.show(), you'll get exactly the output you want:

+-----+---+----------+
|class|val|class_mean|
+-----+---+----------+
|    b|  3|       5.0|
|    b|  7|       5.0|
|    a|  1|       2.0|
|    a|  3|       2.0|
+-----+---+----------+

Method 2: GroupBy + Join

If you prefer a more explicit approach (similar to how you might think about the Python transform() under the hood), you can first compute the grouped means, then join the result back to the original DataFrame.

First, calculate the mean for each class group:

val classMeans = df.groupBy("class").agg(mean($"val").alias("class_mean"))

Then join this aggregated DataFrame back to the original one on the class column:

val resultDf = df.join(classMeans, "class")

This will produce the exact same output as the window function method.

Why This Works

Python's transform() in this scenario computes the mean per group and then "broadcasts" that value to every row in the original group. Both Scala approaches replicate this behavior:

  • Window functions apply the aggregation across all rows in the partition (group), so every row gets the group's mean.
  • The groupBy+join approach computes the mean once per group, then merges that value back with all rows from the original group.

内容的提问来源于stack exchange,提问作者Ladenkov Vladislav

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.21 04:06:39