如何用自定义函数构建ColumnTransformer管道生成目标衍生列?
问题描述
我尝试创建一个ColumnTransformer管道用于通用数据处理,核心需求是生成衍生列(比如立方列)并保留其他原始列。编写的代码如下:
from sklearn.compose import ColumnTransformer from sklearn.preprocessing import FunctionTransformer def take_cube(data,col): data"speed_cube" = np.power(col,3) return data def sin_angle(data, speed, acc): data"angle" = data[speed] * np.sin(data[acc]) return data preprocessor = ColumnTransformer( transformers=[ ("cube", FunctionTransformer(take_cube, validate=False), [speed]), ("sin_angle", FunctionTransformer(sin_angle, kw_args={"speed":"speed","acc":"acceleration"},validate=False), [speed, acceleration]), ], remainder="passthrough" ).set_output(transform="pandas")
运行后生成了大量多余列(如cube__speed、sin_angle__speed等原列副本),但我只需要新增的衍生列(speed_cube、angle)和剩余原始列,请问该如何修改?
解决方案
问题根源在于ColumnTransformer会默认保留转换器输入的所有列,再追加你新增的列。要实现仅生成衍生列的目标,需要调整自定义函数的返回逻辑,同时修正代码中的语法错误:
核心修改点
- 自定义函数不要修改输入数据集,而是返回仅包含衍生列的二维数组/DataFrame,避免污染原始数据并减少冗余列
- 修正语法错误:将中文引号替换为英文引号
" - 补充缺失的
numpy导入 - 转换器的输入列列表会以子数据集形式传入函数,需通过索引或列名正确提取对应列值
修正后的代码
from sklearn.compose import ColumnTransformer from sklearn.preprocessing import FunctionTransformer import numpy as np def take_cube(X): # X为输入的单一列,返回立方计算后的二维数组 return np.power(X, 3).reshape(-1, 1) def sin_angle(X): # X为包含指定列的子DataFrame,通过索引提取列值计算 speed = X.iloc[:, 0] acc = X.iloc[:, 1] return (speed * np.sin(acc)).values.reshape(-1, 1) preprocessor = ColumnTransformer( transformers=[ # 将转换器名称设为衍生列名,最终输出列名会直接用这个名称 ("speed_cube", FunctionTransformer(take_cube, validate=False), ["speed"]), ("angle", FunctionTransformer(sin_angle, validate=False), ["speed", "acceleration"]), ], remainder="passthrough" ).set_output(transform="pandas")
可选优化(用列名直接访问)
如果希望更直观地用列名提取数据,可开启validate=True(确保输入为DataFrame时保留列名):
def sin_angle(X): # 开启validate=True后,X会保留原始列名 return (X["speed"] * np.sin(X["acceleration"])).values.reshape(-1, 1) preprocessor = ColumnTransformer( transformers=[ ("speed_cube", FunctionTransformer(take_cube, validate=True), ["speed"]), ("angle", FunctionTransformer(sin_angle, validate=True), ["speed", "acceleration"]), ], remainder="passthrough" ).set_output(transform="pandas")
关键说明
ColumnTransformer会给每个转换器的输出列加上名称前缀,将转换器名称设为衍生列名即可直接得到目标列名- 自定义函数必须返回二维格式数据(数组或DataFrame),因此需要用
reshape(-1, 1)确保输出符合Scikit-learn要求 - 遵循Scikit-learn无副作用转换原则,不要在函数内修改输入数据集
内容的提问来源于stack exchange,提问作者Chronicles
相关产品推荐
相关产品推荐

