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

求助:PySpark中实现多列目标编码转换的UDF编写

问题:PySpark中实现目标编码UDF,输出同结构DataFrame

需求概述

  • 编写UDF转换PySpark DataFrame,仅对指定列做目标编码映射,输出与原DataFrame结构一致的结果
  • 已使用category_encoders==2.5.1的TargetEncoder训练好XGBoost模型,小数据集用Pandas处理正常,现在需处理大数据集,单独提取目标编码转换结果时遇到问题

可行的预测UDF代码

处理大数据集时,以下预测UDF可正常运行:

from category_encoders import TargetEncoder
te = TargetEncoder(cols = cat_cols)
# (省略模型训练代码)

@pandas_udf('float')
def predict_pandas_udf(*cols, names = features):
    X = pd.concat(cols, axis=1)
    old_names = ["_"+str(x) for x in range(len(names))]
    X.rename(columns=dict(zip(old_names, names)), inplace=True)
    X = te.transform(X)
    return pd.Series(model.predict_proba(X)[:,1])

df = df.withColumn('score', predict_pandas_udf(*df[features]))

无法正常运行的目标编码转换UDF

尝试单独提取目标编码结果时,以下代码无法工作:

@pandas_udf('float')
def target_encoding_udf(df, features):
  X = df[features]
  X = te.transform(X)
  return X

df = df.transform(target_encoding_udf, features)

问题分析与修复方案

问题点

  1. @pandas_udf('float')指定返回单浮点列,但实际需要返回多列DataFrame,类型不匹配
  2. PySpark的pandas_udf用于transform方法时,需定义StructType返回类型,且函数仅接受单个Pandas DataFrame参数
  3. 原UDF的参数写法不符合transform方法对UDF的要求

修复后的代码

步骤1:定义输出Schema

根据原DataFrame结构,将目标编码列替换为浮点类型,其余列保持原类型:

from pyspark.sql.types import StructType, StructField, FloatType
import pandas as pd
from pyspark.sql.functions import pandas_udf

# 获取原DataFrame的schema,复制后替换目标编码列类型
original_schema = df.schema
new_fields = []
for field in original_schema.fields:
    if field.name in cat_cols:  # cat_cols为需做目标编码的列名列表
        new_fields.append(StructField(field.name, FloatType(), nullable=field.nullable))
    else:
        new_fields.append(field)
output_schema = StructType(new_fields)

步骤2:编写正确的Pandas UDF

@pandas_udf(output_schema)
def target_encoding_udf(pdf):
    # 对指定列执行目标编码转换
    encoded_pdf = te.transform(pdf)
    # 确保返回DataFrame的列顺序与原schema一致
    return encoded_pdf[original_schema.fieldNames()]

# 应用UDF获取编码后的DataFrame
df_encoded = df.transform(target_encoding_udf)

关键说明

  • 用pandas_udf(output_schema)指定完整的输出DataFrame结构,而非单个浮点类型
  • transform方法传入的UDF仅接受分区对应的Pandas DataFrame参数(pdf),无需额外传features,直接使用已定义的cat_cols
  • 转换后需保证返回DataFrame的列顺序与原schema一致,避免结构不匹配

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 15:33:30