PySpark自定义异常值检测函数报TypeError的排查求助
PySpark异常值检测函数TypeError排查与修复
问题描述
对PySpark DataFrame执行异常值检测时,使用自定义find_outliers函数,运行new_train = find_outliers(final_matrix)触发TypeError,提示传入的是生成器对象而非字符串或列。该函数在某一DataFrame上可正常运行,但在另一个上报错,已尝试修改数字格式的列名,问题仍存在。
自定义函数代码
def find_outliers(df): # 识别Spark DataFrame中的数值列 numeric_columns = [column[0] for column in df.dtypes if column[1] in ['int','float','long']] # 循环为每个特征创建标记异常值的新列 for column in numeric_columns: less_Q1 = f'less_Q1_{column}' more_Q3 = 'more_Q3_{}'.format(column) Q1 = 'Q1_{}'.format(column) Q3 = 'Q3_{}'.format(column) # 计算四分位数:Q1为第一四分位,Q3为第三四分位 Q1 = df.approxQuantile(column,[0.25],relativeError=0) Q3 = df.approxQuantile(column,[0.75],relativeError=0) # 计算四分位距IQR IQR = Q3[0] - Q1[0] # 定义异常值范围:Q1-1.5*IQR 到 Q3+1.5*IQR less_Q1 = Q1[0] - 1.5*IQR more_Q3 = Q3[0] + 1.5*IQR isOutlierCol = 'is_outlier_{}'.format(column) df = df.withColumn(isOutlierCol,F.when((df[column] > more_Q3) | (df[column] < less_Q1), 1).otherwise(0)) # 筛选出所有标记异常值的列 selected_columns = [column for column in df.columns if column.startswith("is_outlier")] # 将所有异常值标记列求和,生成total_outliers列统计每行异常值数量 df = df.withColumn('total_outliers',sum(df[column] for column in selected_columns)) # 删除临时创建的异常值标记列 df = df.drop(*[column for column in df.columns if column.startswith("is_outlier")]) return df
报错信息
TypeError: Invalid argument, not a string or column: <generator object find_outliers.<locals>.<genexpr> at 0x7f19f4248660> of type <class 'generator'>. For column literals, use 'lit', 'array', 'struct' or 'create_map' function.
错误原因
报错核心源于计算total_outliers的代码行:
df = df.withColumn('total_outliers',sum(df[column] for column in selected_columns))
这里误用了Python内置的sum()函数,它无法处理PySpark的列表达式生成器。
- 当DataFrame只有1个数值列时,
selected_columns仅含1个元素,Python的sum()会直接返回该列对象(刚好符合PySpark要求),因此能正常运行; - 当DataFrame有多个数值列时,生成器包含多个列对象,Python的
sum()无法解析该生成器,最终传入withColumn的是未执行的生成器对象,触发TypeError。
修复方案
需要用PySpark原生函数实现行内多列求和,以下两种方式任选:
方式1:用functools.reduce累加列
from pyspark.sql import functions as F from functools import reduce def find_outliers(df): numeric_columns = [column[0] for column in df.dtypes if column[1] in ['int','float','long']] for column in numeric_columns: less_Q1 = f'less_Q1_{column}' more_Q3 = 'more_Q3_{}'.format(column) Q1 = df.approxQuantile(column,[0.25],relativeError=0) Q3 = df.approxQuantile(column,[0.75],relativeError=0) IQR = Q3[0] - Q1[0] less_Q1 = Q1[0] - 1.5*IQR more_Q3 = Q3[0] + 1.5*IQR isOutlierCol = 'is_outlier_{}'.format(column) df = df.withColumn(isOutlierCol,F.when((df[column] > more_Q3) | (df[column] < less_Q1), 1).otherwise(0)) selected_columns = [column for column in df.columns if column.startswith("is_outlier")] # 修改求和逻辑:用reduce逐个累加所有异常值标记列 if selected_columns: total_outliers_expr = reduce(lambda a, b: a + F.col(b), selected_columns[1:], F.col(selected_columns[0])) df = df.withColumn('total_outliers', total_outliers_expr) else: # 无数值列时,生成值为0的total_outliers列 df = df.withColumn('total_outliers', F.lit(0)) df = df.drop(*[column for column in df.columns if column.startswith("is_outlier")]) return df
方式2:用F.expr拼接加法表达式
from pyspark.sql import functions as F def find_outliers(df): numeric_columns = [column[0] for column in df.dtypes if column[1] in ['int','float','long']] for column in numeric_columns: less_Q1 = f'less_Q1_{column}' more_Q3 = 'more_Q3_{}'.format(column) Q1 = df.approxQuantile(column,[0.25],relativeError=0) Q3 = df.approxQuantile(column,[0.75],relativeError=0) IQR = Q3[0] - Q1[0] less_Q1 = Q1[0] - 1.5*IQR more_Q3 = Q3[0] + 1.5*IQR isOutlierCol = 'is_outlier_{}'.format(column) df = df.withColumn(isOutlierCol,F.when((df[column] > more_Q3) | (df[column] < less_Q1), 1).otherwise(0)) selected_columns = [column for column in df.columns if column.startswith("is_outlier")] # 修改求和逻辑:用字符串拼接加法表达式,通过F.expr执行 if selected_columns: sum_expr = " + ".join(selected_columns) df = df.withColumn('total_outliers', F.expr(sum_expr)) else: df = df.withColumn('total_outliers', F.lit(0)) df = df.drop(*[column for column in df.columns if column.startswith("is_outlier")]) return df
注意事项
- 确保导入
pyspark.sql.functions(通常别名为F) - 处理
selected_columns为空的情况,避免空表达式报错
内容的提问来源于stack exchange,提问作者snigdha mohapatra
相关产品推荐
相关产品推荐

