MLPClassifier部署时特征数不匹配及get_dummies报错求助
解决MLPClassifier部署时特征数不匹配的问题
这个问题我之前帮别人排查过,核心原因是训练阶段和推理阶段的特征预处理逻辑没有完全对齐,导致输入模型的特征数量和模型训练时的不一致。咱们一步步来解决:
问题根源拆解
你在训练时用pd.get_dummies()对整个数据集的类别特征做了编码,生成了固定的25个特征列;但部署时对用户输入单独做get_dummies(),会因为用户输入的类别不全(甚至出现训练数据里没有的新类别),导致生成的特征列数要么少于25,要么多于25(比如你遇到的76列),最终触发特征数不匹配的错误。
解决方案步骤
1. 统一训练与推理的预处理逻辑
不要分开对训练集和用户输入单独做编码,应该把预处理流程标准化,确保两者生成的特征结构完全一致。
2. 方案一:手动对齐特征列(适合快速修复)
训练阶段:保存特征列结构
在训练完特征编码后,把生成的列名保存下来,方便推理时对齐:
import pandas as pd import pickle from sklearn.neural_network import MLPClassifier from sklearn.model_selection import train_test_split # 加载并清洗数据 df = pd.read_csv("nikita_attrition.csv") df = df.drop('StandardHours', axis=1) # 定义需要编码的类别列 categorical_cols = ['Education','JobSatisfaction', 'WorkLifeBalance', 'EnvironmentSatisfaction', 'StockOptionLevel', 'JobInvolvement', 'PerformanceRating', 'JobLevel'] # 做独热编码 features_in = pd.get_dummies(df, columns=categorical_cols, dummy_na=False) # 保存训练时的特征列名 train_feature_cols = features_in.columns.tolist() with open('train_feature_cols.pkl', 'wb') as f: pickle.dump(train_feature_cols, f) # 准备标签并训练模型 y = pd.get_dummies(df['Attrition'], drop_first=True).values.ravel() X_train, X_test, y_train, y_test = train_test_split(features_in, y, test_size=0.2, random_state=42) model = MLPClassifier() model.fit(X_train, y_train) # 保存模型 with open('mlp_attrition_model.pkl', 'wb') as f: pickle.dump(model, f)
推理阶段:对齐特征列
获取用户输入后,按照训练时的列结构调整特征:
import pandas as pd import pickle # 加载保存的列名和模型 with open('train_feature_cols.pkl', 'rb') as f: train_feature_cols = pickle.load(f) with open('mlp_attrition_model.pkl', 'rb') as f: model = pickle.load(f) def preprocess_user_input(user_input_df): # 移除无关列(和训练时保持一致) user_input_df = user_input_df.drop('StandardHours', axis=1, errors='ignore') # 对用户输入做独热编码 categorical_cols = ['Education','JobSatisfaction', 'WorkLifeBalance', 'EnvironmentSatisfaction', 'StockOptionLevel', 'JobInvolvement', 'PerformanceRating', 'JobLevel'] user_features = pd.get_dummies(user_input_df, columns=categorical_cols, dummy_na=False) # 对齐训练时的特征列: # 1. 添加训练时有但用户输入没有的列,值设为0 for col in train_feature_cols: if col not in user_features.columns: user_features[col] = 0 # 2. 删除训练时没有的列(比如用户输入了新类别生成的列) user_features = user_features[train_feature_cols] return user_features # 示例用户输入(根据你的实际字段调整) user_input = pd.DataFrame([{ 'Education': 2, 'JobSatisfaction': 3, 'WorkLifeBalance': 2, # 其他必填字段... }]) # 预处理后预测 processed_input = preprocess_user_input(user_input) prediction = model.predict(processed_input)
3. 方案二:用Sklearn Pipeline更规范(推荐)
用Pipeline把预处理和模型打包在一起,避免手动对齐列的麻烦,还能减少出错概率:
import pandas as pd import pickle from sklearn.pipeline import Pipeline from sklearn.compose import ColumnTransformer from sklearn.preprocessing import OneHotEncoder from sklearn.neural_network import MLPClassifier from sklearn.model_selection import train_test_split # 加载数据 df = pd.read_csv("nikita_attrition.csv") df = df.drop('StandardHours', axis=1) # 定义类别列 categorical_cols = ['Education','JobSatisfaction', 'WorkLifeBalance', 'EnvironmentSatisfaction', 'StockOptionLevel', 'JobInvolvement', 'PerformanceRating', 'JobLevel'] X_raw = df.drop('Attrition', axis=1) y = pd.get_dummies(df['Attrition'], drop_first=True).values.ravel() # 构建预处理管道 preprocessor = ColumnTransformer( transformers=[ ('cat', OneHotEncoder(sparse_output=False), categorical_cols) ], remainder='passthrough' # 保留数值列不做处理 ) # 构建完整Pipeline pipeline = Pipeline([ ('preprocessor', preprocessor), ('classifier', MLPClassifier()) ]) # 训练 X_train, X_test, y_train, y_test = train_test_split(X_raw, y, test_size=0.2, random_state=42) pipeline.fit(X_train, y_train) # 保存整个Pipeline with open('attrition_pipeline.pkl', 'wb') as f: pickle.dump(pipeline, f)
推理时直接用Pipeline处理:
import pandas as pd import pickle # 加载Pipeline loaded_pipeline = pickle.load(open('attrition_pipeline.pkl', 'rb')) # 用户输入DataFrame user_input = pd.DataFrame([{ 'Education': 2, 'JobSatisfaction': 3, # 其他字段... }]) # 直接预测,Pipeline会自动处理特征对齐 prediction = loaded_pipeline.predict(user_input)
额外注意事项
- 处理缺失值:如果用户可能遗漏某些字段,预处理时要提前用训练集的众数/均值填充,避免编码出错。
- 校验输入类别:如果用户输入了训练数据中没有的类别(比如
Education出现了5,但训练时只有1-4),可以在预处理前做校验,提示用户输入合法值。
内容的提问来源于stack exchange,提问作者Aditi c.
相关产品推荐
相关产品推荐

