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

PySpark合并两个Map数组列:补全缺失键生成单列

问题:合并Spark DataFrame中Array of Maps类型列的正确实现

数据集定义

test_table = spark.createDataFrame(
    [
        ("US", "CA", "S", "2022-10-01",100, 10, 1),
        ("US", "CA", "M", "2022-10-01",100, 15, 5),
        ("US", "CA", "L", "2022-10-01",100, 20, 10),
        ("US", "CA", "S", "2022-10-01",200, 10, 1),
        ("US", "CA", "M", "2022-10-01",200, 15, 5),
        ("US", "CA", "L", "2022-10-01",200, 20, 10),
        ("US", "CA", "S", "2022-10-02",100, 11, 1),
        ("US", "CA", "M", "2022-10-02",100, 13, 3),
        ("US", "CA", "L", "2022-10-02",100, 17, 7),
        ("US", "CA", "S", "2022-10-02",200, 11, 1),
        ("US", "CA", "M", "2022-10-02",200, 13, 3),
    ],
    schema=["country_code","state_code","size","dt","store_id","ttl_sold","ttl_returned"]
)

聚合操作后的DataFrame

执行以下代码后,得到包含latest_payload和prev_payload两个Array of Maps类型列的DataFrame:

w = Window.partitionBy("country_code", "state_code", "size", "store_id").orderBy("dt").rangeBetween(Window.unboundedPreceding,0)
w2 = Window.partitionBy("country_code", "state_code", "size").orderBy("dt")
df_w_cumulative_sum = (
    test_table
    .withColumn("cumulative_ttl_sold", F.sum("ttl_sold").over(w))
    .withColumn("cumulative_ttl_returned", F.sum("ttl_returned").over(w))
    .groupBy("dt","country_code", "state_code", "size")
    .agg(F.collect_list(F.create_map(F.col("store_id"), F.struct(F.col("cumulative_ttl_sold"), F.col("cumulative_ttl_returned")))).alias("latest_payload"))
    .withColumn("prev_payload", F.lag(F.col("latest_payload"), 1).over(w2))
    .where(F.col("dt") == "2022-10-02")
)

数据样例

rowdtcountry_codestate_codesizelatest_payloadprev_payload
12022-10-01USCAL[{"100":{"cumulative_ttl_sold":20,"cumulative_ttl_returned":10}},{"200":{"cumulative_ttl_sold":20,"cumulative_ttl_returned":10}}]null
22022-10-01USCAM[{"100":{"cumulative_ttl_sold":15,"cumulative_ttl_returned":5}},{"200":{"cumulative_ttl_sold":15,"cumulative_ttl_returned":5}}]null
32022-10-01USCAS[{"100":{"cumulative_ttl_sold":10,"cumulative_ttl_returned":1}},{"200":{"cumulative_ttl_sold":10,"cumulative_ttl_returned":1}}]null
42022-10-02USCAL[{"100":{"cumulative_ttl_sold":37,"cumulative_ttl_returned":17}}][{"100":{"cumulative_ttl_sold":20,"cumulative_ttl_returned":10}},{"200":{"cumulative_ttl_sold":20,"cumulative_ttl_returned":10}}]
52022-10-02USCAM[{"100":{"cumulative_ttl_sold":28,"cumulative_ttl_returned":8}},{"200":{"cumulative_ttl_sold":28,"cumulative_ttl_returned":8}}][{"100":{"cumulative_ttl_sold":15,"cumulative_ttl_returned":5}},{"200":{"cumulative_ttl_sold":15,"cumulative_ttl_returned":5}}]
62022-10-02USCAS[{"100":{"cumulative_ttl_sold":21,"cumulative_ttl_returned":2}},{"200":{"cumulative_ttl_sold":21,"cumulative_ttl_returned":2}}][{"100":{"cumulative_ttl_sold":10,"cumulative_ttl_returned":1}},{"200":{"cumulative_ttl_sold":10,"cumulative_ttl_returned":1}}]

合并需求

需要将latest_payload和prev_payload合并为一个Map列,规则为:

  • 保留latest_payload中的所有键值对
  • 补充prev_payload中latest_payload缺失的键值对

以第4行为例,预期输出为:

{'100': {'cumulative_ttl_sold': 37, 'cumulative_ttl_returned': 17}, '200': {'cumulative_ttl_sold': 20, 'cumulative_ttl_returned': 10}}

错误的UDF实现

尝试的UDF返回错误结果:

@F.udf(
    MapType(
        IntegerType(), 
        StructType([
            StructField("cumulative_ttl_sold", LongType(), False), 
            StructField("cumulative_ttl_sold", LongType(), False)
        ])
    )
)
def merge_payloads(lastest_payload, prev_payload):
    payload: Dict[int, Dict[str, int]] = {}
    if prev_payload is not None:
        for latest in lastest_payload:
            for k,v in latest.items():
                payload[k] = v
        for prev in prev_payload:
            for k,v in prev.items():
                if k not in payload.keys():
                    payload[k]=v
                else:
                    break
    else:
        for latest in lastest_payload:
            for k, v in latest.items():
                payload[k] = v
    return payload

问题分析与正确实现

原UDF的问题

  1. Struct字段重复:返回类型定义中,两个字段都命名为cumulative_ttl_sold,第二个字段应为cumulative_ttl_returned
  2. 循环逻辑错误:处理prev_payload时,遇到已存在的key就执行break,导致后续缺失的key无法被补充(比如第4行的200会被跳过)

正确的UDF实现

from pyspark.sql.types import MapType, IntegerType, StructType, StructField, LongType
from pyspark.sql import functions as F
from typing import Dict, List, Optional

@F.udf(
    MapType(
        IntegerType(), 
        StructType([
            StructField("cumulative_ttl_sold", LongType(), False), 
            StructField("cumulative_ttl_returned", LongType(), False)
        ])
    )
)
def merge_payloads(latest_payload: Optional[List[Dict[int, Dict]]], prev_payload: Optional[List[Dict[int, Dict]]]) -> Dict[int, Dict]:
    payload = {}
    # 优先处理latest_payload,保留最新数据
    if latest_payload:
        for item in latest_payload:
            payload.update(item)
    # 处理prev_payload,补充缺失的key
    if prev_payload:
        for item in prev_payload:
            for k, v in item.items():
                if k not in payload:
                    payload[k] = v
    return payload

验证效果

调用UDF查看结果:

result_df = df_w_cumulative_sum.withColumn("merged_payload", merge_payloads(F.col("latest_payload"), F.col("prev_payload")))
result_df.select("size", "merged_payload").show(truncate=False)

第4行(size=L)的merged_payload会得到预期结果:

{'100': Row(cumulative_ttl_sold=37, cumulative_ttl_returned=17), '200': Row(cumulative_ttl_sold=20, cumulative_ttl_returned=10)}

可选:无需UDF的Spark内置函数实现

如果想避免Python UDF的序列化开销,可直接用Spark内置函数完成合并:

from pyspark.sql import functions as F

# 将Array of Maps转为单个Map
latest_map = F.map_from_entries(F.flatten(F.col("latest_payload")))
prev_map = F.map_from_entries(F.flatten(F.col("prev_payload")))
# 合并两个Map,优先保留latest_map的键值对
merged_map = F.map_concat(latest_map, F.map_filter(prev_map, lambda k, v: F.not(F.map_contains(latest_map, k))))

result_df = df_w_cumulative_sum.withColumn("merged_payload", merged_map)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.14 01:05:19