如何从sklearn pipeline输出预处理后的数据对象以传入FLAML训练
问题解答
首先明确结论:完全可以直接从Sklearn预处理管线输出预处理后的数据,传入FLAML进行训练。
现有代码的错误点
你写的pp_training_data, pp_training_label = preprocessor_pipeline属于错误用法:
- ColumnTransformer/Pipeline本身是预处理规则容器,没有执行拟合、转换操作,无法直接拆分为特征和标签
- 预处理仅作用于特征数据,标签不需要经过预处理转换
修改后的可运行代码
import pandas as pd from sklearn.model_selection import train_test_split from sklearn.pipeline import Pipeline from sklearn.impute import SimpleImputer from sklearn.preprocessing import StandardScaler, OneHotEncoder from sklearn.compose import ColumnTransformer def diamond_preprocess(data_dir): data = pd.read_csv(data_dir) cleaned_data = data.drop(['id', 'depth_percent'], axis=1) # 移除不需要的特征 x = cleaned_data.drop(['price'], axis=1) # 特征集 y = cleaned_data['price'] # 标签集 x_train, x_test, y_train, y_test = train_test_split(x, y, test_size=0.2, random_state=42) numerical_features = x_train.select_dtypes(include=['int64', 'float64']).columns categorical_features = x_train.select_dtypes(include=['object']).columns numerical_transformer = Pipeline(steps=[ ('imputer', SimpleImputer(strategy='median')), # 缺失值用中位数填充 ('scaler', StandardScaler()) # 数值特征标准化 ]) categorical_transformer = Pipeline(steps=[ ('imputer', SimpleImputer(strategy='constant', fill_value='missing')), # 缺失值用固定值填充 ('onehot', OneHotEncoder(handle_unknown='ignore')) # 分类特征独热编码 ]) preprocessor_pipeline = ColumnTransformer( transformers=[ ('num', numerical_transformer, numerical_features), ('cat', categorical_transformer, categorical_features) ]) # 仅在训练集拟合预处理规则,避免数据泄露 preprocessor_pipeline.fit(x_train) # 转换得到预处理后的训练集、测试集特征 pp_training_data = preprocessor_pipeline.transform(x_train) pp_test_data = preprocessor_pipeline.transform(x_test) # 返回预处理后的训练特征、训练标签、测试特征、测试标签,按需取用即可 return pp_training_data, y_train, pp_test_data, y_test
后续调用FLAML的方式
你只需要拿到返回的预处理后数据,直接调用fit方法即可:
pp_training_data, pp_training_labels, _, _ = diamond_preprocess("你的数据路径.csv") automl.fit(X_train=pp_training_data, y_train=pp_training_labels, **automl_settings)
额外提示
如果需要将预处理后的numpy数组转回DataFrame便于排查问题,可以通过preprocessor_pipeline.get_feature_names_out()方法获取处理后的所有特征名,再和数组拼接即可。
内容的提问来源于stack exchange,提问作者Luleo_Primoc
相关产品推荐
相关产品推荐

