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

SKLearn pipeline使用独热编码时fit与predict特征数不匹配问题求解

问题根因

你出现特征维度不匹配的核心原因是自定义转换器不符合Scikit-Learn的设计规范:fit阶段没有记录训练集的统计信息,transform阶段直接使用当前传入数据的统计特征做处理,导致训练和预测阶段生成的特征不一致。
具体到你的代码有两个核心问题:

  • handleCategory类直接在transform阶段调用pd.get_dummies,没有记录训练集所有类别型特征的取值范围,测试集类别少于训练集时,生成的独热列数自然更少
  • handleImputation类在transform阶段才计算缺失率决定要删除的列、计算填充用的均值/众数,不仅会导致特征列数可能不一致,还会产生数据泄露问题

可行解决方案

修改后的自定义转换器代码

import pandas as pd
from sklearn.base import BaseEstimator, TransformerMixin

class HandleCategory(BaseEstimator, TransformerMixin):
    def __init__(self):
        self.train_columns = None  # 存储训练集独热编码后的所有列名
        self.category_cols = None  # 存储所有类别型列名
        self.cat_unique_values = {}  # 存储每个类别列的所有唯一取值
    
    def fit(self, X, y=None):
        X_ = X.copy()
        # 先筛选所有类别型列
        self.category_cols = X_.select_dtypes(include=['object', 'category']).columns.tolist()
        # 记录每个类别列的唯一值
        for col in self.category_cols:
            self.cat_unique_values[col] = X_[col].unique().tolist()
        # 生成训练集独热编码后的列名,作为后续transform的标准
        dummy_train = pd.get_dummies(X_, columns=self.category_cols)
        self.train_columns = dummy_train.columns.tolist()
        return self
    
    def transform(self, X, y=None):
        X_ = X.copy()
        # 按照训练集的类别取值做独热编码,避免测试集缺失类别导致列少
        for col in self.category_cols:
            X_[col] = pd.Categorical(X_[col], categories=self.cat_unique_values[col])
        dummy_test = pd.get_dummies(X_, columns=self.category_cols)
        # 补全训练集有但测试集没有的列,赋值为0
        for col in self.train_columns:
            if col not in dummy_test.columns:
                dummy_test[col] = 0
        # 按照训练集列顺序重排,保证特征顺序一致
        dummy_test = dummy_test[self.train_columns]
        return dummy_test

class HandleImputation(BaseEstimator, TransformerMixin):
    def __init__(self, missing_threshold=0.7):
        self.missing_threshold = missing_threshold
        self.drop_cols = []  # 存储训练集确定要删除的高缺失率列
        self.fill_values = {}  # 存储每个列的填充值(均值/众数)
        self.has_missing_cols = []  # 存储需要加missing标记的列
    
    def fit(self, X, y=None):
        X_ = X.copy()
        data_len = X_.shape[0]
        # 训练阶段就确定要删除的高缺失率列
        self.drop_cols = [col for col in X_.columns if X_[col].isnull().sum() > data_len * self.missing_threshold]
        # 剩下的列计算填充值
        remain_cols = [col for col in X_.columns if col not in self.drop_cols]
        for col in remain_cols:
            missing_cnt = X_[col].isnull().sum()
            if missing_cnt > 0:
                self.has_missing_cols.append(col)
                # 类别型用众数填充,数值型用均值填充
                if X_[col].dtype == "object":
                    self.fill_values[col] = X_[col].mode().iloc[0]
                else:
                    self.fill_values[col] = X_[col].mean()
        return self
    
    def transform(self, X, y=None):
        X_ = X.copy()
        # 先删除训练集确定要删的列
        X_ = X_.drop(columns=self.drop_cols, errors='ignore')
        # 加缺失标记+填充缺失值
        for col in self.has_missing_cols:
            if col in X_.columns:
                X_[f'{col}_missing'] = X_[col].isnull()
                X_[col] = X_[col].fillna(self.fill_values[col])
        return X_

Pipeline组装调用代码

from sklearn.pipeline import Pipeline
from sklearn.ensemble import GradientBoostingRegressor

gbr_params = {
    'n_estimators': 1000,
    'max_depth': 3,
    'min_samples_split': 5,
    'learning_rate': 0.01,
    'loss': 'ls'
}
gbr = GradientBoostingRegressor(**gbr_params)

# 组装pipeline,注意处理顺序:缺失值处理 -> 类别编码 -> 模型
pipeline1 = Pipeline(steps=[
    ('imputation', HandleImputation()),
    ('category_encode', HandleCategory()),
    ('model', gbr)
])

# 正常训练预测即可
pipeline1.fit(train_data, trainY)
testY = pipeline1.predict(test_data)

额外优化建议

如果不想自定义转换器,也可以直接用Scikit-Learn自带的OneHotEncoder(设置handle_unknown='ignore'即可自动忽略测试集未见过的类别,保证特征数一致)、SimpleImputer、ColumnTransformer组合实现相同功能,原生接口兼容性更好。

内容的提问来源于stack exchange,提问作者Prabu M

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 16:45:04