在sklearn pipeline中对指定列应用KBinsDiscretizer返回DataFrame报错求助
问题原因
报错触发是因为ColumnTransformer的目标列传入了单个字符串'numeric_col_to_change',sklearn会将该列作为1D的Series传入后续的KBinsDiscretizer,而该转换器要求输入必须为2D结构(哪怕仅有1列也需要是(n_samples, 1)的格式)。
另外当前的自定义PandasColumnTransformer存在两个隐藏问题:
- 未设置
remainder='passthrough',会默认丢弃未指定变换的列 - 输出列名直接复用原输入的列名,若列顺序发生变化会导致列名错位
修复方案
- 将单列选择改为列表格式,避免传入1D数据
- 给
PandasColumnTransformer添加remainder='passthrough'参数保留原始列 - 优化
PandasColumnTransformer的列名生成逻辑,适配列顺序变化
完整可运行代码
import pandas as pd from sklearn.compose import ColumnTransformer from sklearn.preprocessing import KBinsDiscretizer from sklearn.pipeline import Pipeline class PandasColumnTransformer(ColumnTransformer): def get_feature_names_out(self, input_features=None): feature_names = [] for name, trans, cols, _ in self._iter(fitted=True): if trans == 'drop': continue elif trans == 'passthrough': feature_names.extend(cols) else: feature_names.extend(trans.get_feature_names_out(cols)) return feature_names def transform(self, X: pd.DataFrame) -> pd.DataFrame: return pd.DataFrame(super().transform(X), columns=self.get_feature_names_out(), index=X.index) def fit_transform(self, X: pd.DataFrame, y=None) -> pd.DataFrame: return pd.DataFrame(super().fit_transform(X, y), columns=self.get_feature_names_out(), index=X.index) class PandasKBinsDiscretizer(KBinsDiscretizer): def __init__(self, n_bins): super().__init__(n_bins, encode='ordinal') def get_feature_names_out(self, input_features=None): return input_features def transform(self, X): self.col_names = list(X.columns.values) X = super().transform(X) return pd.DataFrame(X, columns=self.col_names) binner_on_numeric = PandasColumnTransformer( transformers=[ ("binner", PandasKBinsDiscretizer(2), ['numeric_col_to_change']) ], remainder='passthrough' # 保留未指定变换的列 ) pp = Pipeline([('binner_just_numeric', binner_on_numeric)]) d = {'numeric_col_not_to_change': [1, 2, 1, 2, 1, 2], 'numeric_col_to_change': [1, 2, 3, 4, 5, 6]} df = pd.DataFrame(data=d) res = pp.fit_transform(df) assert isinstance(res, pd.DataFrame)
验证结果
运行后输出的res为标准pandas DataFrame格式,numeric_col_to_change会被转换为0/1的分箱结果,其余列保持原始值,assert语句不会触发报错。
内容的提问来源于stack exchange,提问作者Yehoshaphat Schellekens
相关产品推荐
相关产品推荐

