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

自定义sklearn TransformerMixin:保障训练测试数据转换一致性

解决方案:让自定义转换器的transform输出与fit阶段完全一致

你的问题核心在于原转换器没有在fit阶段保存训练数据的关键特征信息,导致transform测试数据时完全依赖当前数据的结构,从而出现列不匹配的问题。scikit-learn的fit/transform设计就是为了在fit阶段固化训练数据的统计特征,测试阶段严格复用这些特征来保证一致性,下面是修改后的完整方案:

核心思路

在fit方法中记录每一列的关键信息:

  • 分类列(object类型)的所有唯一类别
  • 数值列(int64类型)的分箱边界
    在transform方法中,严格使用这些保存的信息转换测试数据,确保:
  • 分类列只生成训练时存在的类别列,测试数据中的新类别直接忽略(或标记为NaN)
  • 数值列使用训练时的分箱边界,不管测试数据的数值范围如何,分箱规则和训练时一致
  • 最终输出的列顺序、数量完全和fit_transform的结果匹配,缺失列自动补0

修改后的完整代码

import numpy as np
import pandas as pd
from sklearn.base import TransformerMixin, BaseEstimator  # 加上BaseEstimator更规范
from sklearn.linear_model import LogisticRegression

class DFTransformer(BaseEstimator, TransformerMixin):
    def fit(self, df, y=None, **fit_params):
        # 初始化保存训练信息的字典
        self.col_types = {}  # 记录每列的类型
        self.categories = {}  # 记录分类列的所有类别
        self.bins = {}  # 记录数值列的分箱边界
        self.feature_names = []  # 记录最终输出的所有特征列名

        for col in df.columns:
            dtype = df[col].dtype.name
            self.col_types[col] = dtype
            
            if dtype == 'object':
                # 保存该分类列的所有唯一类别,排序保证列顺序稳定
                self.categories[col] = sorted(df[col].unique())
                # 生成对应的dummy列名
                dummy_cols = [f"{col}_{cat}" for cat in self.categories[col]]
                self.feature_names.extend(dummy_cols)
            elif dtype == 'int64':
                # 计算分箱边界并保存,这里保持原代码的bins=5
                s = df[col].copy()
                _, bins = pd.cut(s, bins=5, retbins=True)
                # 修正bins的边界,避免测试数据超出范围时出现NaN
                bins[0] = -np.inf
                bins[-1] = np.inf
                self.bins[col] = bins
                # 生成对应的分箱dummy列名
                cut_labels = pd.cut(df[col], bins=self.bins[col]).cat.categories
                dummy_cols = [f"{col}_{label}" for label in cut_labels]
                self.feature_names.extend(dummy_cols)
        
        # 去重并固定列顺序(防止重复列,虽然原逻辑不会有,但保险)
        self.feature_names = list(dict.fromkeys(self.feature_names))
        return self

    def transform(self, df, **trans_params):
        self.STACKER = pd.DataFrame(index=df.index, columns=self.feature_names).fillna(0)
        
        for col in df.columns:
            if col not in self.col_types:
                # 测试数据出现训练时没有的列,直接忽略
                continue
                
            dtype = self.col_types[col]
            if dtype == 'object':
                # 将测试数据转换为指定类别的Categorical,新类别会变成NaN
                cat_series = pd.Categorical(df[col], categories=self.categories[col])
                # 生成dummy列,只保留训练时的类别
                dummy_df = pd.get_dummies(cat_series, prefix=col)
                # 将dummy_df的值更新到STACKER对应的列
                for dummy_col in dummy_df.columns:
                    if dummy_col in self.STACKER.columns:
                        self.STACKER[dummy_col] = dummy_df[dummy_col]
            elif dtype == 'int64':
                # 使用训练时的分箱边界进行分箱
                cut_series = pd.cut(df[col], bins=self.bins[col])
                # 生成dummy列
                dummy_df = pd.get_dummies(cut_series, prefix=col)
                # 更新到STACKER
                for dummy_col in dummy_df.columns:
                    if dummy_col in self.STACKER.columns:
                        self.STACKER[dummy_col] = dummy_df[dummy_col]
        
        return self.STACKER

    # 可选:添加get_feature_names_out方法,符合scikit-learn最新API规范
    def get_feature_names_out(self, input_features=None):
        return self.feature_names

代码解释

  1. fit阶段的关键改进:

    • 保存了每一列的类型、分类列的所有类别、数值列的分箱边界,以及最终输出的特征列名列表
    • 对数值列的分箱边界做了优化:把第一个边界设为-inf,最后一个设为inf,避免测试数据超出训练数据范围时出现NaN(也可以根据需求调整)
    • 固定了特征列的顺序,保证输出稳定
  2. transform阶段的关键改进:

    • 先创建一个和训练时完全一样的空DataFrame,所有列初始化为0
    • 分类列转换时,使用pd.Categorical强制指定训练时的类别,测试数据中的新类别会被转为NaN,不会生成新列
    • 数值列转换时,严格使用训练时的分箱边界,不管测试数据的数值范围如何,分箱规则和训练时一致
    • 只更新训练时存在的列,测试数据多余的列直接忽略

测试验证

用你提供的测试代码验证:

# 模拟训练数据
df = pd.DataFrame({'integers': np.random.randint(2000, 20000, 30, dtype='int64'), 
                   'categorical': np.random.choice(list('ABCDEFGHIJKLMNOP'), 30)}, 
                  columns=['integers', 'categorical'])
trans = DFTransformer()
X = trans.fit_transform(df)
y = np.random.binomial(1, 0.5, 30)
lr = LogisticRegression()
lr.fit(X, y)

# 模拟测试数据(包含训练时没有的类别和更大的数值范围)
X_test = pd.DataFrame({'integers': np.random.randint(2000, 60000, 30, dtype='int64'), 
                       'categorical': np.random.choice(list('ABGIOPXYZ'), 30)}, 
                      columns=['integers', 'categorical'])

# 现在转换测试数据并预测,不会再报错
X_test_transformed = trans.transform(X_test)
print(f"训练数据特征数:{X.shape[1]},测试数据特征数:{X_test_transformed.shape[1]}")  # 应该相等
predictions = lr.predict(X_test_transformed)
print(predictions)

额外建议

  • 继承BaseEstimator可以让你的转换器兼容scikit-learn的很多工具(比如Pipeline、GridSearchCV等)
  • 可以添加参数控制新类别的处理方式(比如是否把新类别映射到一个统一的"其他"列)
  • 对于数值列的分箱,也可以考虑使用分位数分箱(pd.qcut),这样分箱边界更稳定,不会因为训练数据的极端值影响
  • 可以添加对NaN值的处理逻辑(比如在fit阶段记录每列的填充值,transform时统一填充)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 08:43:55