基于部门类型移除PySpark DataFrame中的异常值
移除PySpark DataFrame各部门年龄异常值的正确做法
我们有两个PySpark DataFrame:
df:存储员工的部门、年龄、薪资等基础信息joined_read:按部门聚合计算出的年龄统计数据,结构如下:
| 部门(Department) | 平均年龄(Mean_Age) | 标准差(Std_dev) | 年龄上限(upperbound) | 年龄下限(lowerbound) |
|---|---|---|---|---|
| Dept1 | 38 | 7 | 52 | 24 |
| Dept2 | 43 | 4 | 51 | 35 |
| Dept3 | 32 | 6 | 44 | 22 |
需求是根据每个部门的年龄上下界,批量移除所有部门的年龄异常值,得到过滤后的单一DataFrame。之前尝试用循环遍历部门列表标记异常值但未得到有效结果,尝试的代码如下:
dept_list=(joined_read.select('Dept').distinct().rdd.map(lambda x : x[0]).collect()) for dept in dept_list: upper_bound=joined_read.filter(F.col('Dept')==dept).select('upperbound').collect()[0][0] lower_bound=joined_read.filter(F.col('Dept')==dept).select('lowerbound').collect()[0][0] filt_df=df.withColumn('outlier',when((F.col('Age')>=lower_bound)&\ (F.col('Age')<=upper_bound)&\ (F.col('Dept')==dept),"0"))\ .otherwise("1"))
原代码问题分析
- 循环覆盖数据:每次循环都基于原始
df生成新的filt_df,前一次部门的异常值标记会被后一次覆盖,最终只有最后一个部门的标记生效。 - 性能与风险问题:多次调用
collect()会把分布式数据拉到Driver端,数据量大时容易触发内存溢出,完全违背Spark分布式计算的设计逻辑。
正确解决方案
直接用Spark的关联操作,将两个DataFrame按部门匹配,一次性获取每个员工对应的年龄上下界,再完成过滤,无需循环:
from pyspark.sql import functions as F # 按部门关联,给每个员工匹配对应部门的年龄上下界 df_with_bounds = df.join( joined_read.select("Department", "upperbound", "lowerbound"), on="Department", how="inner" ) # 过滤年龄不在上下界内的异常值,最后移除临时关联的上下界列 filtered_df = df_with_bounds.filter( (F.col("Age") >= F.col("lowerbound")) & (F.col("Age") <= F.col("upperbound")) ).drop("upperbound", "lowerbound") # 查看过滤后的结果 filtered_df.show()
如果需要保留异常值标记(而非直接删除),可以这样修改:
df_with_bounds = df.join( joined_read.select("Department", "upperbound", "lowerbound"), on="Department", how="left" ) # 新增outlier列标记是否为异常值 df_with_outlier = df_with_bounds.withColumn( "outlier", F.when( (F.col("Age") >= F.col("lowerbound")) & (F.col("Age") <= F.col("upperbound")), "0" ).otherwise("1") ).drop("upperbound", "lowerbound")
内容的提问来源于stack exchange,提问作者Amit Lohani
相关产品推荐
相关产品推荐

