如何使用PySpark计算各部门员工的薪资中位数?
在PySpark中实现分组薪资中位数计算(无内置函数)
确实,PySpark SQL并没有提供直接的中位数聚合函数,不过我们可以通过两种主流思路来实现按部门分组计算薪资中位数的需求,下面我会一步步拆解实现过程,同时兼顾小数据和大数据场景的不同方案。
核心思路概述
我们的目标是按Department分组后计算每组的薪资中位数,核心逻辑是:
- 小数据场景:分组收集薪资列表 → 排序列表 → 根据列表长度的奇偶性计算中位数
- 大数据场景:利用窗口函数计算分位数 → 筛选分位值接近0.5的记录取平均(避免内存压力)
方案一:基于collect_list + UDF的小数据实现
这个方案适合分组数据量不大的场景,代码直观易懂:
1. 准备测试DataFrame
首先我们创建一个模拟的员工薪资数据集:
from pyspark.sql import SparkSession from pyspark.sql.functions import collect_list, udf from pyspark.sql.types import DoubleType # 初始化Spark会话 spark = SparkSession.builder.appName("DeptSalaryMedian").getOrCreate() # 构造测试数据 sample_data = [ ("HR", 101, 5000), ("HR", 102, 6500), ("HR", 103, 7200), ("Engineering", 201, 8000), ("Engineering", 202, 9500), ("Engineering", 203, 10000), ("Engineering", 204, 12000), ("Finance", 301, 7000) ] df = spark.createDataFrame(sample_data, ["Department", "Employee ID", "Salary"]) df.show()
2. 定义中位数计算UDF
因为collect_list返回的列表是无序的,我们需要先排序,再根据列表长度计算中位数:
def compute_median(salary_list): # 对薪资列表进行排序 sorted_salaries = sorted(salary_list) count = len(sorted_salaries) if count % 2 == 1: # 奇数个元素,直接取中间位置的值 return sorted_salaries[count // 2] else: # 偶数个元素,取中间两个值的平均值 return (sorted_salaries[count // 2 - 1] + sorted_salaries[count // 2]) / 2 # 注册UDF,指定返回类型为DoubleType median_udf = udf(compute_median, DoubleType())
3. 分组计算中位数
通过groupBy分组后收集薪资列表,再调用UDF计算中位数:
from pyspark.sql.functions import col result_df = df.groupBy("Department") \ .agg(collect_list(col("Salary")).alias("salary_collection")) \ .withColumn("Median_Salary", median_udf(col("salary_collection"))) \ .select("Department", "Median_Salary") result_df.show()
方案二:基于窗口函数的大数据实现
如果你的数据集非常大,collect_list会将整个分组的数据拉到单个节点处理,容易出现内存溢出。这时可以用窗口函数的分位数逻辑来实现,全程分布式处理:
from pyspark.sql.window import Window from pyspark.sql.functions import percent_rank, avg # 定义窗口:按部门分区,薪资升序排序 dept_window = Window.partitionBy("Department").orderBy("Salary") # 计算每条薪资记录在所属部门的百分比排名 ranked_df = df.withColumn("percent_rank", percent_rank().over(dept_window)) # 筛选排名接近0.5的记录(容错区间0.49-0.51,避免浮点精度问题),然后取平均得到中位数 big_data_median_result = ranked_df.filter(col("percent_rank").between(0.49, 0.51)) \ .groupBy("Department") \ .agg(avg("Salary").alias("Median_Salary")) big_data_median_result.show()
这个方案不需要收集整个分组的数据,更适合生产环境的大数据场景。
内容的提问来源于stack exchange,提问作者Tran Dinh Cuong
相关产品推荐
相关产品推荐

