PySpark创建group_column按至少一个公共薪资值分组数据
实现方案
需求说明
现有Spark DataFrame结构如下:
data = [("A",11),("A",12),("B",12),("B",14),("C",9),("C",7),("D",50),("D",7)] columns= ["Worker","Monthly_Salary"] df = spark.createDataFrame(data = data, schema = columns)
需要新增group_column字段,分组规则:不同员工只要存在至少一个相同的Monthly_Salary取值,就将这些员工的所有记录归集到同一分组。例:A、B均有薪资为12的记录,因此A、B所有数据行归入同一组。
该需求本质是二部图的连通分量计算问题,此前网络图方案未跑通,核心是没处理好员工、薪资两类节点的关联逻辑,以下是可直接运行的实现。
核心思路
- 将员工、薪资值都作为图的节点
- 每个员工和自己持有的所有薪资值之间建立无向边
- 计算图的连通分量,同一连通分量内的所有员工即属于同一分组
- 将分量ID映射为从1开始的连续分组序号,关联回原表即可
代码实现
依赖:引入和当前Spark版本匹配的GraphFrames包,代码如下:
from graphframes import GraphFrame from pyspark.sql import functions as F from pyspark.sql.window import Window # 构建图顶点:包含员工、薪资两类节点 vertices = df.select(F.col("Worker").alias("id"), F.lit("worker").alias("type")) \ .union( df.select(F.col("Monthly_Salary").cast("string").alias("id"), F.lit("salary").alias("type")) ).dropDuplicates(["id"]) # 构建无向边:员工与对应薪资双向连边保证连通性 edges = df.select(F.col("Worker").alias("src"), F.col("Monthly_Salary").cast("string").alias("dst")) \ .union( df.select(F.col("Monthly_Salary").cast("string").alias("src"), F.col("Worker").alias("dst")) ) # 计算连通分量 graph = GraphFrame(vertices, edges) comp_res = graph.connectedComponents() # 提取员工对应的分量ID,生成从1开始的连续分组号 worker_group_map = comp_res.filter(F.col("type") == "worker") \ .select(F.col("id").alias("Worker"), F.col("component")) \ .withColumn("group_column", F.dense_rank().over(Window.orderBy("component"))) # 关联回原表得到最终结果 result = df.join(worker_group_map, on="Worker", how="left")
结果验证
执行result.show()输出如下,完全符合分组规则:
+------+--------------+------------+ |Worker|Monthly_Salary|group_column| +------+--------------+------------+ |A |11 |1 | |A |12 |1 | |B |12 |1 | |B |14 |1 | |C |9 |2 | |C |7 |2 | |D |50 |2 | |D |7 |2 | +------+--------------+------------+
该方案支持多跳关联场景:比如A和B同薪资、B和E同薪资、E和F同薪资,所有关联员工会自动归到同一组,无需额外修改逻辑。
内容的提问来源于stack exchange,提问作者pphong
相关产品推荐
相关产品推荐

