Sklearn自定义ColumnTransformer管道训练报TypeError问题
问题原因
两个核心问题导致报错:
- 数值特征处理流水线顺序错误:
num_transformer中先执行StandardScaler标准化,后执行SimpleImputer缺失值填补。StandardScaler本身不支持NaN输入,低版本sklearn遇到含NaN的输入时内部方法会返回None,触发解包报错;高版本sklearn会直接抛出输入包含NaN的显式错误。缺失值填补逻辑必须放在所有缩放、编码类操作之前。 - 自定义特征构造器不符合sklearn接口规范:
Add_family类的__init__方法入参名为add_family,但赋值给实例属性时错写为self.ad_family,入参名和实例属性名不一致会导致sklearn克隆估计器、参数搜索、模型持久化时出现异常。
修正后的代码
import pandas as pd import numpy as np from sklearn.base import BaseEstimator, TransformerMixin from sklearn.pipeline import Pipeline from sklearn.preprocessing import StandardScaler, OneHotEncoder from sklearn.impute import SimpleImputer from sklearn.compose import ColumnTransformer, make_column_selector from sklearn.linear_model import LogisticRegression # 工具函数提到类外部,避免每次调用transform都重复定义 def get_family_type(var): if var == 1: return 'alone' elif var <= 4: return 'small' else: return 'big' class Add_family(BaseEstimator, TransformerMixin): def __init__(self, add_family = True): # 修正:实例属性名和__init__入参名保持一致 self.add_family = add_family def fit(self, X, y=None): return self def transform(self, X, y=None): df = pd.DataFrame(X).copy() if self.add_family: df['Family_size'] = df.apply(lambda x: x.Parch + x.SibSp + 1, axis=1) df['FamilyType'] = df.apply(lambda x: get_family_type(x.Family_size), axis=1) df = df.drop(columns=['Parch', 'SibSp']) return df # 修正:缺失值填补放在标准化之前 num_transformer = Pipeline([ ('imputer', SimpleImputer(strategy='mean')), ('scaler', StandardScaler()) ]) cat_transformer = Pipeline([ ('onehot', OneHotEncoder(handle_unknown='ignore')) ]) col_transform = ColumnTransformer([ ('cat', cat_transformer, make_column_selector(dtype_include=object)), ('num', num_transformer, make_column_selector(dtype_include=np.number)) ]) preprocessor = Pipeline([ ('Adder_features', Add_family(add_family=True)), ('transform', col_transform) ]) lr = Pipeline([ ('prep', preprocessor), ('clf', LogisticRegression()) ]) # 执行训练 features = ['Pclass', 'Sex', 'Age', 'Parch', 'SibSp','Embarked'] X_train = data[features] y_train = data['Survived'] lr.fit(X_train, y_train)
验证说明
修正后单独运行预处理器、嵌套入完整训练流水线均可正常执行,不会触发类型错误。
内容的提问来源于stack exchange,提问作者Сергей Клочко
相关产品推荐
相关产品推荐

