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") )
数据样例
| row | dt | country_code | state_code | size | latest_payload | prev_payload |
|---|---|---|---|---|---|---|
| 1 | 2022-10-01 | US | CA | L | [{"100":{"cumulative_ttl_sold":20,"cumulative_ttl_returned":10}},{"200":{"cumulative_ttl_sold":20,"cumulative_ttl_returned":10}}] | null |
| 2 | 2022-10-01 | US | CA | M | [{"100":{"cumulative_ttl_sold":15,"cumulative_ttl_returned":5}},{"200":{"cumulative_ttl_sold":15,"cumulative_ttl_returned":5}}] | null |
| 3 | 2022-10-01 | US | CA | S | [{"100":{"cumulative_ttl_sold":10,"cumulative_ttl_returned":1}},{"200":{"cumulative_ttl_sold":10,"cumulative_ttl_returned":1}}] | null |
| 4 | 2022-10-02 | US | CA | L | [{"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}}] |
| 5 | 2022-10-02 | US | CA | M | [{"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}}] |
| 6 | 2022-10-02 | US | CA | S | [{"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的问题
- Struct字段重复:返回类型定义中,两个字段都命名为
cumulative_ttl_sold,第二个字段应为cumulative_ttl_returned - 循环逻辑错误:处理
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
相关产品推荐
相关产品推荐

