Spark 2.2中使用Scala无需explode实现聚合值的方法
explode in Spark 2.2 (Scala) Got it, let's tackle this problem step by step. Since you're working with Spark 2.2 and want to avoid explode (which can cause data bloat with large arrays), we'll use User-Defined Functions (UDFs) to compute aggregations directly within each array entry first, then roll up the results at the group level. Here's how to do it:
Step 1: Define the Schema & Load Data
First, we need to explicitly define the schema to handle the nested array of structs (especially since some entries have empty structs). Then load your JSON data into a DataFrame:
import org.apache.spark.sql.SparkSession import org.apache.spark.sql.functions._ import org.apache.spark.sql.types._ // Initialize Spark Session val spark = SparkSession.builder() .appName("AggWithoutExplode") .master("local[*]") // Remove this line for cluster mode .getOrCreate() import spark.implicits._ // Raw JSON input data val jsonInput = """{"amount":"2.00","cal_group":[{}],"set_id":7057} {"amount":"1.00","cal_group":[{}],"set_id":7057} {"amount":"7.00","cal_group": [{"abc_cd":"abc00160","abc_cnt":6.0,"cde_cnt":7.0},{"abc_cd":"abc00160","abc_cnt":5.0,"cde_cnt":2.0},{"abc_cd":"abc00249","abc_cnt":0.0,"cde_cnt":1.0}],"set_id":7057}""" // Define schema for the nested `cal_group` array val calGroupStruct = StructType(Seq( StructField("abc_cd", StringType, nullable = true), StructField("abc_cnt", DoubleType, nullable = true), StructField("cde_cnt", DoubleType, nullable = true) )) // Define full DataFrame schema val dfSchema = StructType(Seq( StructField("amount", StringType, nullable = true), StructField("cal_group", ArrayType(calGroupStruct), nullable = true), StructField("set_id", IntegerType, nullable = true) )) // Load data into DataFrame val rawDf = spark.read.schema(dfSchema).json(spark.sparkContext.parallelize(jsonInput.split("\n"))) // Convert `amount` from String to Double for numeric aggregation val typedDf = rawDf.withColumn("amount", $"amount".cast(DoubleType))
Step 2: Create a UDF to Aggregate Within Each Array
Spark 2.2 doesn't support the higher-order aggregate function (introduced in 2.3), so we'll write a UDF to iterate over each cal_group array and sum up the abc_cnt and cde_cnt values:
// UDF to calculate sum of abc_cnt and cde_cnt for a single row's cal_group array val sumCalGroupUdf = udf((calGroupEntries: Seq[Row]) => { calGroupEntries.foldLeft((0.0, 0.0)) { case ((totalAbc, totalCde), entry) => // Handle null values by defaulting to 0.0 val abc = entry.getAs[Option[Double]]("abc_cnt").getOrElse(0.0) val cde = entry.getAs[Option[Double]]("cde_cnt").getOrElse(0.0) (totalAbc + abc, totalCde + cde) } }) // Apply the UDF to get per-row aggregations val rowLevelAggDf = typedDf .withColumn("cal_group_sums", sumCalGroupUdf($"cal_group")) .select( $"set_id", $"amount", $"cal_group_sums._1".alias("row_total_abc"), $"cal_group_sums._2".alias("row_total_cde") )
Step 3: Group by set_id and Compute Final Aggregations
Now we can group by set_id and sum up the per-row values to get the final aggregated results:
val finalAggDf = rowLevelAggDf.groupBy($"set_id") .agg( sum($"amount").alias("total_amount"), sum($"row_total_abc").alias("total_abc_cnt"), sum($"row_total_cde").alias("total_cde_cnt") ) // Show the result finalAggDf.show()
Expected Output
+------+------------+-------------+-------------+ |set_id|total_amount|total_abc_cnt|total_cde_cnt| +------+------------+-------------+-------------+ | 7057| 10.0| 11.0| 10.0| +------+------------+-------------+-------------+
Why This Works
- We avoid
explodeby handling array aggregation at the row level first, which prevents data duplication and keeps performance better for large datasets. - The UDF safely handles empty structs and null values by defaulting missing counts to 0.0.
内容的提问来源于stack exchange,提问作者Vijay_Shinde

