You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在PySpark中按组用均值与5倍标准差移除多列异常值

Multi-Column Outlier Detection with Spark (5σ Rule)

Got it, let's tackle this problem step by step. You want to extend your existing single-column outlier logic to handle both price and income, group by cd/segment, and generate binary flags (1 = outlier to remove, 0 = keep) for each column. Here's how to refactor your code to be scalable and clean:

First, Let's Recap the Problem

We have a Spark DataFrame and need to:

  • Group rows by cd and segment
  • For both price and income, mark values as outliers if they fall outside mean ± 5*stddev
  • Create a dedicated flag column (is_outlier_<col>) for each metric, where 1 means outlier, 0 means keep

Original Data & Single-Column Code

First, let's restate your sample data and existing implementation for context:

Sample Data

from pyspark.sql import SparkSession
import pyspark.sql.functions as f

spark = SparkSession.builder.appName("OutlierDetection").getOrCreate()

data = [ ('a', '1',20,10), ('a', '1',30,16), ('a', '1',50,91), ('a', '1',60,34), ('a', '1',200,23), 
         ('a', '2',33,87), ('a', '2',86,90), ('a','2',89,35), ('a', '2',90,24), ('a', '2',40,97), ('a', '2',1,21), 
         ('b', '1',45,96), ('b', '1',56,99), ('b', '1',89,23), ('b', '1',98,64), 
         ('b', '2',86,42), ('b', '2',45,54), ('b', '2',67,95), ('b','2',86,70), ('b', '2',91,64), ('b', '2',2,53), ('b', '2',4,87) ]
df = spark.createDataFrame(data, ['cd','segment','price','income'])

Existing Single-Column Code (Only Handles price)

mean_std = ( df 
             .groupBy('cd', 'segment') 
             .agg( *[f.mean(colName).alias(f'mean_{colName}') for colName in ['price']], 
                   *[f.stddev(colName).alias(f'stddev_{colName}') for colName in ['price']]) )

mean_columns = ['mean_price']
std_columns = ['stddev_price']

upper = mean_std
for col_1 in mean_columns:
    for col_2 in std_columns:
        if col_1 != col_2:
            name = f'{col_1}_upper_limit'
            upper = upper.withColumn(name, f.col(col_1) + f.col(col_2)*5)

lower = upper
for col_1 in mean_columns:
    for col_2 in std_columns:
        if col_1 != col_2:
            name = f'{col_1}_lower_limit'
            lower = lower.withColumn(name, f.col(col_1) - f.col(col_2)*5)

outliers = (df.join(lower, how = 'left', on = ['cd', 'segment']) 
            .withColumn('is_outlier_price', 
                        f.when((f.col('price')>f.col('mean_price_upper_limit')) | (f.col('price')<f.col('mean_price_lower_limit')),1) 
                        .otherwise(0)) )  # Changed None to 0 to match your requirement

Optimized Multi-Column Solution

Instead of duplicating code for each column, we can generalize the logic to handle any number of target columns with a simple loop. This makes the code easier to maintain if you add more metrics later:

# Define the columns we want to check for outliers
target_cols = ['price', 'income']

# Step 1: Calculate mean and stddev for each target column, grouped by cd/segment
mean_std_df = df.groupBy('cd', 'segment').agg(
    *[f.mean(col).alias(f'mean_{col}') for col in target_cols],
    *[f.stddev(col).alias(f'stddev_{col}') for col in target_cols]
)

# Step 2: Compute upper and lower bounds (mean ± 5*stddev) for each column
bounds_df = mean_std_df
for col in target_cols:
    bounds_df = bounds_df.withColumn(
        f'{col}_upper',
        f.col(f'mean_{col}') + 5 * f.col(f'stddev_{col}')
    ).withColumn(
        f'{col}_lower',
        f.col(f'mean_{col}') - 5 * f.col(f'stddev_{col}')
    )

# Step 3: Join bounds back to original data and create outlier flags
final_df = df.join(bounds_df, on=['cd', 'segment'], how='left')

for col in target_cols:
    final_df = final_df.withColumn(
        f'is_outlier_{col}',
        f.when(
            (f.col(col) > f.col(f'{col}_upper')) | (f.col(col) < f.col(f'{col}_lower')),
            1
        ).otherwise(0)
    )

# Optional: Drop intermediate columns (mean, stddev, bounds) if you don't need them
final_df = final_df.drop(*[col for col in bounds_df.columns if col not in ['cd', 'segment']])

# View the result
final_df.show(truncate=False)

What's Better About This Approach?

  • Scalable: Just add more column names to target_cols and the code handles them automatically—no need to rewrite logic for each metric.
  • Cleaner: Replaced nested loops with a single loop over target columns, making the code easier to read and debug.
  • Consistent: Ensures the same outlier rule applies to all columns, and strictly uses 1/0 for flags (no None values, which aligns with your requirement).
  • Flexible: You can keep or drop the intermediate mean/stddev/bound columns as needed.

内容的提问来源于stack exchange,提问作者Gun

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.08 09:57:29