如何用PySpark或SQL基于Tan匹配值分组聚合客户编码与Tan列表?
需求说明
现有如下结构的DataFrame,其中CustomerCode为4-1234和4-1235的记录共享Tan值MUMS12345A。需要将所有共享至少一个Tan值的CustomerCode合并为列表,同时将对应所有Tan值去重后合并为列表。
输入数据
| CustomerCode | TanList |
|---|---|
| 4-1234 | MUMS12345A,BLRS12345E,BLRS12345G |
| 4-1235 | MUMS12345A,CHED12345A |
| 4-1236 | RTKD12345A |
目标输出
| CustomerCodeList | TANList |
|---|---|
| 4-1234, 4-1235 | MUMS12345A,BLRS12345E,BLRS12345G,CHED12345A |
| 4-1236 | RTKD12345A |
PySpark 实现方案
该需求属于**连通分量(Connected Components)**问题:共享Tan值的客户属于同一个连通组,可通过以下步骤实现:
- 拆分
TanList列,将每个Tan值拆分为单独行 - 建立客户与Tan的关联,通过Tan关联到其他客户形成连通组
- 为每个连通组分配唯一标识
- 按组聚合客户列表和去重后的Tan列表
基础实现(无额外依赖)
from pyspark.sql import SparkSession from pyspark.sql.functions import split, explode, collect_set, concat_ws from pyspark.sql.window import Window # 初始化SparkSession spark = SparkSession.builder.appName("CustomerTanGrouping").getOrCreate() # 创建输入DataFrame data = [ ("4-1234", "MUMS12345A,BLRS12345E,BLRS12345G"), ("4-1235", "MUMS12345A,CHED12345A"), ("4-1236", "RTKD12345A") ] df = spark.createDataFrame(data, ["CustomerCode", "TanList"]) # 1. 拆分TanList为单独行 df_exploded = df.withColumn("tan", explode(split(df.TanList, ","))) # 2. 获取每个Tan对应的所有客户 tan_customer_map = df_exploded.groupBy("tan").agg(collect_set("CustomerCode").alias("linked_customers")) # 3. 迭代合并连通组,直到组不再变化 def merge_components(input_df): while True: # 合并每个客户关联的所有客户组 merged = input_df.join(tan_customer_map, on="tan", how="inner")\ .groupBy("CustomerCode")\ .agg(collect_set("linked_customers").alias("all_groups"))\ .withColumn("flattened", explode("all_groups"))\ .groupBy("CustomerCode")\ .agg(collect_set("flattened").alias("final_group")) # 检查是否完成合并 prev_group_count = input_df.select("CustomerCode").distinct().count() input_df = input_df.join(merged, on="CustomerCode", how="inner").drop("tan") current_group_count = input_df.select("final_group").distinct().count() if prev_group_count == current_group_count: break return input_df df_connected = merge_components(df_exploded) # 4. 按组聚合生成最终结果 final_df = df_connected.join(df_exploded, on="CustomerCode", how="inner")\ .groupBy("final_group")\ .agg( concat_ws(", ", collect_set("CustomerCode")).alias("CustomerCodeList"), concat_ws(", ", collect_set("tan")).alias("TANList") )\ .drop("final_group") # 展示结果 final_df.show(truncate=False)
高效实现(使用GraphFrames)
若环境支持安装GraphFrames,可直接用图算法计算连通分量:
from graphframes import GraphFrame # 构建顶点:客户作为顶点 vertices = df.select("CustomerCode").distinct().withColumnRenamed("CustomerCode", "id") # 构建边:Tan作为中间节点,连接共享Tan的客户 edges1 = df_exploded.select("CustomerCode", "tan").withColumnRenamed("CustomerCode", "src").withColumnRenamed("tan", "dst") edges2 = df_exploded.select("tan", "CustomerCode").withColumnRenamed("tan", "src").withColumnRenamed("CustomerCode", "dst") edges = edges1.union(edges2) # 创建图并计算连通分量 g = GraphFrame(vertices, edges) connected_df = g.connectedComponents() # 按连通组聚合 result = connected_df.join(df_exploded, connected_df.id == df_exploded.CustomerCode, how="inner")\ .groupBy("component")\ .agg( concat_ws(", ", collect_set("id")).alias("CustomerCodeList"), concat_ws(", ", collect_set("tan")).alias("TANList") )\ .drop("component") result.show(truncate=False)
SQL 实现方案
以Spark SQL为例,思路与PySpark一致,通过递归CTE处理连通分量:
-- 创建临时表 CREATE TEMPORARY TABLE customer_tan ( CustomerCode STRING, TanList STRING ); -- 插入数据 INSERT INTO customer_tan VALUES ('4-1234', 'MUMS12345A,BLRS12345E,BLRS12345G'), ('4-1235', 'MUMS12345A,CHED12345A'), ('4-1236', 'RTKD12345A'); WITH exploded_tan AS ( -- 拆分TanList为单独行 SELECT CustomerCode, explode(split(TanList, ',')) AS tan FROM customer_tan ), tan_customer_map AS ( -- 获取每个Tan对应的所有客户 SELECT tan, collect_set(CustomerCode) AS customer_group FROM exploded_tan GROUP BY tan ), recursive_components AS ( -- 递归合并连通组 SELECT CustomerCode AS customer, collect_set(CustomerCode) AS component FROM exploded_tan GROUP BY CustomerCode UNION ALL SELECT rc.customer, collect_set(DISTINCT elem) AS component FROM recursive_components rc JOIN exploded_tan et ON EXISTS (SELECT 1 FROM unnest(rc.component) r WHERE r = et.CustomerCode) JOIN tan_customer_map tcm ON et.tan = tcm.tan, unnest(tcm.customer_group) elem GROUP BY rc.customer ), final_components AS ( -- 去重并保留每个客户的最终连通组 SELECT customer, component FROM recursive_components WHERE array_length(component) = ( SELECT max(array_length(component)) FROM recursive_components rc2 WHERE rc2.customer = recursive_components.customer ) GROUP BY customer, component ) -- 按组聚合生成最终输出 SELECT concat_ws(', ', fc.component) AS CustomerCodeList, concat_ws(', ', collect_set(et.tan)) AS TANList FROM final_components fc JOIN exploded_tan et ON et.CustomerCode = fc.customer GROUP BY fc.component;
内容的提问来源于stack exchange,提问作者scriptbees 22
相关产品推荐
相关产品推荐

