scikit-learn中独热编码后多变量插补及Pipeline构建方案问询
解决方案:结合分类列辅助的多变量插补Pipeline实现
需求回顾
处理包含分类列与数值列的数据集,完成以下操作并整合到标准Pipeline中:
- 对分类列执行OneHot编码
- 利用编码后的分类列辅助数值列的缺失值插补(使用
IterativeImputer) - 最终输出保留原始分类列与插补后的数值列
完整代码实现
import pandas as pd from sklearn.pipeline import Pipeline from sklearn.preprocessing import OneHotEncoder from sklearn.impute import IterativeImputer from sklearn.base import BaseEstimator, TransformerMixin # 自定义转换器:利用分类列辅助插补数值列 class ImputeWithCategorical(BaseEstimator, TransformerMixin): def __init__(self, numeric_col, categorical_cols): self.numeric_col = numeric_col self.categorical_cols = categorical_cols self.onehot = OneHotEncoder(sparse_output=False, drop='first') self.imputer = IterativeImputer() def fit(self, X, y=None): # 拟合分类列的OneHot编码器 self.onehot.fit(X[self.categorical_cols]) # 拼接数值列与编码后的分类列,用于拟合插补器 encoded_cats = self.onehot.transform(X[self.categorical_cols]) fit_dataset = pd.concat([X[self.numeric_col], pd.DataFrame(encoded_cats)], axis=1) self.imputer.fit(fit_dataset) return self def transform(self, X): # 对分类列执行编码 encoded_cats = self.onehot.transform(X[self.categorical_cols]) # 拼接数值列与编码后的分类列,执行插补 transform_dataset = pd.concat([X[self.numeric_col], pd.DataFrame(encoded_cats)], axis=1) imputed_results = self.imputer.transform(transform_dataset) # 替换原始DataFrame中的数值列,保留分类列 output = X.copy() output[self.numeric_col] = imputed_results[:, 0] return output # 示例数据集 sample_data = pd.DataFrame({ "a": [4.4, 1.0, None, 3.0, 2.7], "b": ["HIGH", "HIGH", "LOW", "HIGH", "LOW"], "c": [True, False, False, True, False] }) # 构建Pipeline pipeline = Pipeline(steps=[ ('impute_with_categorical', ImputeWithCategorical(numeric_col='a', categorical_cols=['b', 'c'])) ]) # 拟合并转换数据 imputed_data = pipeline.fit_transform(sample_data) print(imputed_data)
代码说明
- 自定义转换器
ImputeWithCategorical:fit方法:先完成分类列的OneHot编码拟合,再将数值列与编码后的分类特征拼接,让IterativeImputer学习到特征间的关联,为后续插补做准备。transform方法:对输入数据的分类列编码,拼接数值列后执行插补,最后将插补后的数值列替换回原始数据集,保留分类列的原始格式。
- Pipeline整合:将整个流程封装为标准的scikit-learn Pipeline,具备
fit和transform方法,可直接用于训练和预测流程,也能与其他scikit-learn组件兼容。
内容的提问来源于stack exchange,提问作者rhn89
相关产品推荐
相关产品推荐

