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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.25 20:36:17