One-Hot编码后训练与测试DataFrame列数不匹配,如何对齐?
解决One-Hot编码后训练集与测试集列对齐问题
我有训练集和测试集两个初始列数相同的DataFrame,但分类数据列中二者取值不同,经过One-Hot编码后列数不一致,无法进行预测。试过pd.get_dummies和ColumnTransformer,以下是之前的编码代码片段:
for k in range(80): if isinstance(X_train.iloc[0,k],str): lista0.append(X_train.columns[k]) indice0.append(k) else: lista1.append(X_train.columns[k]) indice1.append(k) for z in range(80): if isinstance(X_test.iloc[0,z],str): lista2.append(X_test.columns[z]) indice2.append(z) else: lista3.append(X_test.columns[z]) indice3.append(z) X_train_tran = ColumnTransformer([('onehot',OneHotEncoder(sparse_output=False),indice0),('nothing','passthrough',indice1)]) X_test_tran = ColumnTransformer([('onehot',OneHotEncoder(sparse_output=False),indice2),('nothing','passthrough',indice3)]) X1_train = X_train_tran.fit_transform(X_train) X1_test = X_test_tran.fit_transform(X_test)
问题根源
之前的操作错误在于分别对训练集和测试集独立拟合编码器,导致两个OneHotEncoder各自学习了不同的类别集合,最终编码后的列数、列名完全不匹配。正确逻辑是:只在训练集上拟合编码器,再用同一个编码器转换测试集。
正确实现步骤
1. 统一识别分类列(无需重复处理测试集)
因为初始列数一致,直接从训练集识别分类/数值列即可:
lista0 = [] # 分类列名称 lista1 = [] # 数值列名称 indice0 = [] # 分类列索引 indice1 = [] # 数值列索引 for idx, col in enumerate(X_train.columns): if isinstance(X_train.iloc[0, idx], str): lista0.append(col) indice0.append(idx) else: lista1.append(col) indice1.append(idx)
2. 用训练集拟合转换器,复用转换测试集
只在训练集执行fit_transform,测试集直接调用transform:
from sklearn.compose import ColumnTransformer from sklearn.preprocessing import OneHotEncoder # 定义预处理转换器 preprocessor = ColumnTransformer( [('onehot', OneHotEncoder(sparse_output=False, handle_unknown='ignore'), indice0), ('passthrough', 'passthrough', indice1)] ) # 训练集:拟合+转换 X1_train = preprocessor.fit_transform(X_train) # 测试集:直接转换(复用训练集学到的类别规则) X1_test = preprocessor.transform(X_test)
3. 可选:转回DataFrame并对齐列名
如果需要将数组转回DataFrame方便查看/后续操作,可以提取完整特征名:
import pandas as pd # 获取One-Hot编码后的特征名 onehot_features = preprocessor.named_transformers_['onehot'].get_feature_names_out(lista0) # 合并所有特征名 all_features = list(onehot_features) + lista1 # 转换为DataFrame X1_train_df = pd.DataFrame(X1_train, columns=all_features) X1_test_df = pd.DataFrame(X1_test, columns=all_features)
关键注意点
- 绝对不能对测试集单独调用
fit或fit_transform,必须复用训练集的编码器规则。 handle_unknown='ignore'参数可处理测试集出现训练集没有的类别(自动忽略对应列),如果希望报错提醒,可设置为handle_unknown='error'。
内容的提问来源于stack exchange,提问作者Angelo
相关产品推荐
相关产品推荐

