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

scikit-learn中独热编码后多变量插补及Pipeline构建方案问询

解决方案:结合分类列辅助的多变量插补Pipeline实现

需求回顾

处理包含分类列与数值列的数据集,完成以下操作并整合到标准Pipeline中:

  1. 对分类列执行OneHot编码
  2. 利用编码后的分类列辅助数值列的缺失值插补(使用IterativeImputer)
  3. 最终输出保留原始分类列与插补后的数值列

完整代码实现

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 02:27:45