Spark中如何用UDF合并不同结构体类型的数组字段?
PySpark 实现方案
实现思路
- 把
arrayOne和arrayTwo分别转成以Q为键的字典,通过Q值快速定位元素,避免低效的循环匹配 - 收集两个数组中所有唯一的Q值,确保不遗漏任何元素
- 对每个Q值:
- 若
arrayOne存在对应元素,保留其a/b/c/Q字段,将x/y/z设为None - 若仅
arrayTwo存在对应元素,保留其x/y/z/Q字段,将a/b/c设为None
- 若
- 把生成的结构体列表转换为Spark支持的数组类型
代码实现
from pyspark.sql import SparkSession from pyspark.sql.functions import udf from pyspark.sql.types import StructType, StructField, StringType, ArrayType # 定义输出结构体类型 output_struct = StructType([ StructField("a", StringType(), nullable=True), StructField("b", StringType(), nullable=True), StructField("c", StringType(), nullable=True), StructField("Q", StringType(), nullable=False), StructField("x", StringType(), nullable=True), StructField("y", StringType(), nullable=True), StructField("z", StringType(), nullable=True) ]) def merge_arrays(array_one, array_two): # 转换为Q键映射 map_one = {item["Q"]: item for item in array_one} if array_one else {} map_two = {item["Q"]: item for item in array_two} if array_two else {} # 收集所有Q值 all_qs = set(map_one.keys()).union(set(map_two.keys())) result = [] for q in all_qs: if q in map_one: merged = { "a": map_one[q]["a"], "b": map_one[q]["b"], "c": map_one[q]["c"], "Q": q, "x": None, "y": None, "z": None } else: merged = { "a": None, "b": None, "c": None, "Q": q, "x": map_two[q]["x"], "y": map_two[q]["y"], "z": map_two[q]["z"] } result.append(merged) return result # 注册UDF merge_arrays_udf = udf(merge_arrays, ArrayType(output_struct)) # 测试示例 spark = SparkSession.builder.appName("MergeArrays").getOrCreate() # 构造测试数据 data = [ ( [{"a": "a1", "b": "b1", "c": "c1", "Q": "q1"}, {"a": "a2", "b": "b2", "c": "c2", "Q": "q2"}], [{"x": "x2", "y": "y2", "z": "z2", "Q": "q2"}, {"x": "x3", "y": "y3", "z": "z3", "Q": "q3"}] ) ] schema = StructType([ StructField("arrayOne", ArrayType(StructType([ StructField("a", StringType()), StructField("b", StringType()), StructField("c", StringType()), StructField("Q", StringType()) ]))), StructField("arrayTwo", ArrayType(StructType([ StructField("x", StringType()), StructField("y", StringType()), StructField("z", StringType()), StructField("Q", StringType()) ]))) ]) df = spark.createDataFrame(data, schema) df.withColumn("arrayThree", merge_arrays_udf("arrayOne", "arrayTwo")).show(truncate=False)
Scala Spark 实现方案
实现思路
和PySpark逻辑一致,通过将数组转换为Map[String, Row](以Q为键),合并所有Q键后生成目标数组,完全避免explode_outer带来的性能损耗和数据结构破坏。
代码实现
import org.apache.spark.sql.SparkSession import org.apache.spark.sql.functions.udf import org.apache.spark.sql.types._ object MergeArraysUDF { def main(args: Array[String]): Unit = { val spark = SparkSession.builder.appName("MergeArrays").getOrCreate() import spark.implicits._ // 定义输出结构体类型 val outputStruct = StructType(Seq( StructField("a", StringType, nullable = true), StructField("b", StringType, nullable = true), StructField("c", StringType, nullable = true), StructField("Q", StringType, nullable = false), StructField("x", StringType, nullable = true), StructField("y", StringType, nullable = true), StructField("z", StringType, nullable = true) )) // 核心合并逻辑 def mergeArrays(arrayOne: Seq[Row], arrayTwo: Seq[Row]): Seq[Row] = { val mapOne = if (arrayOne != null) arrayOne.map(row => row.getAs[String]("Q") -> row).toMap else Map.empty[String, Row] val mapTwo = if (arrayTwo != null) arrayTwo.map(row => row.getAs[String]("Q") -> row).toMap else Map.empty[String, Row] // 收集所有唯一Q值 val allQs = mapOne.keys ++ mapTwo.keys allQs.map { q => mapOne.get(q) match { case Some(row) => // 优先取arrayOne元素,x/y/z设为null Row( row.getAs[String]("a"), row.getAs[String]("b"), row.getAs[String]("c"), q, null.asInstanceOf[String], null.asInstanceOf[String], null.asInstanceOf[String] ) case None => // 取arrayTwo元素,a/b/c设为null val row = mapTwo(q) Row( null.asInstanceOf[String], null.asInstanceOf[String], null.asInstanceOf[String], q, row.getAs[String]("x"), row.getAs[String]("y"), row.getAs[String]("z") ) } }.toSeq } // 注册UDF val mergeArraysUdf = udf(mergeArrays _, ArrayType(outputStruct)) // 构造测试数据 val data = Seq( ( Seq( Row("a1", "b1", "c1", "q1"), Row("a2", "b2", "c2", "q2") ), Seq( Row("x2", "y2", "z2", "q2"), Row("x3", "y3", "z3", "q3") ) ) ) val schema = StructType(Seq( StructField("arrayOne", ArrayType(StructType(Seq( StructField("a", StringType), StructField("b", StringType), StructField("c", StringType), StructField("Q", StringType) )))), StructField("arrayTwo", ArrayType(StructType(Seq( StructField("x", StringType), StructField("y", StringType), StructField("z", StringType), StructField("Q", StringType) )))) )) val df = spark.createDataFrame(data, schema) df.withColumn("arrayThree", mergeArraysUdf($"arrayOne", $"arrayTwo")).show(truncate = false) } }
内容的提问来源于stack exchange,提问作者Gligorijevic
相关产品推荐
相关产品推荐

