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

Sklearn Pipeline中获取OneHot编码列名报错及特征列名提取需求

解决Sklearn Pipeline中获取OneHot编码列名的问题

你的错误原因很明确:pipe['preprocessor'].transformers[0][0]取到的是分类预处理器的名称字符串'categorical',而不是实际的OneHotEncoder实例,自然没法调用get_feature_names方法。

下面是正确的操作步骤:

1. 务必先拟合Pipeline

只有拟合后,OneHotEncoder才能根据训练数据生成对应的类别列名,所以先执行拟合:

# 假设X是你的输入特征数据,y是目标变量
pipe.fit(X, y)

2. 获取OneHotEncoder实例并生成编码列名

推荐用named_transformers_属性通过名称索引,比直接取列表索引更可靠:

# 获取拟合后的OneHotEncoder实例
ohe = pipe['preprocessor'].named_transformers_['categorical']
# 生成编码后的列名(Sklearn 1.0+版本推荐用get_feature_names_out,替代旧的get_feature_names)
encoded_cols = ohe.get_feature_names_out(categorical_columns)

如果你的Sklearn版本较低(<1.0),可以用get_feature_names:

encoded_cols = ohe.get_feature_names(categorical_columns)

3. 获取完整的特征列名(含passthrough的列)

如果需要Pipeline处理后所有特征的列名(包括分类编码列+原非分类列),可以直接调用ColumnTransformer的get_feature_names_out方法:

all_feature_names = pipe['preprocessor'].get_feature_names_out()

完整示例代码

import pandas as pd
from sklearn.compose import ColumnTransformer
from sklearn.preprocessing import OneHotEncoder
from sklearn.ensemble import RandomForestRegressor
from sklearn.pipeline import make_pipeline

# 模拟数据
data = pd.DataFrame({
    'cat_col1': ['A', 'B', 'C', 'A'],
    'cat_col2': ['X', 'Y', 'X', 'Y'],
    'num_col1': [10, 20, 30, 40],
    'target': [100, 200, 300, 400]
})
categorical_columns = ['cat_col1', 'cat_col2']
X = data.drop('target', axis=1)
y = data['target']

# 构建Pipeline
categorical_preprocessor = OneHotEncoder(handle_unknown="ignore")
preprocessor = ColumnTransformer(
    [('categorical', categorical_preprocessor, categorical_columns)], remainder="passthrough")
est = RandomForestRegressor(n_estimators=100, random_state=0)
pipe = make_pipeline(preprocessor, est)

# 拟合Pipeline
pipe.fit(X, y)

# 获取OneHot编码后的列名
ohe = pipe['preprocessor'].named_transformers_['categorical']
encoded_cols = ohe.get_feature_names_out(categorical_columns)
print("OneHot编码列名:", encoded_cols)

# 获取所有处理后的特征列名
all_cols = pipe['preprocessor'].get_feature_names_out()
print("所有特征列名:", all_cols)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.03 14:05:22