Flask集成预训练RandomForestClassifier时特征数量不匹配报错求助
Flask集成RandomForestClassifier特征不匹配问题解决
错误信息
ValueError: X has 3 features, but DecisionTreeClassifier is expecting 35 features as input
问题背景
开发集成预训练RandomForestClassifier的Flask应用时,用划分数据集得到的X_test测试模型完全正常,但用户通过HTML表单提交的数据测试时,出现上述特征数量不匹配错误。已尝试将表单数据存入字典转成pd.DataFrame,再转成csr_matrix格式,问题仍未解决。
核心问题
你在预测路由里重新初始化OneHotEncoder和ColumnTransformer并调用了fit_transform方法——这会让新编码器仅根据当前用户的3条数据生成特征,而非沿用训练模型时基于全量训练集生成的35维特征规则。两者特征维度不一致,自然触发模型输入不匹配的错误。
解决方案
1. 训练阶段保存预处理编码器
训练模型时,必须把用于特征编码的ColumnTransformer(及内部的OneHotEncoder)和模型一起保存,示例代码:
from joblib import dump, load from sklearn.compose import ColumnTransformer from sklearn.preprocessing import OneHotEncoder from sklearn.ensemble import RandomForestClassifier # 训练时的特征编码逻辑 categorical_features = ["DFOS","CP","PASSION"] one_hot = OneHotEncoder(handle_unknown='ignore') # 兼容训练集未出现的类别 transformer = ColumnTransformer( [("one_hot", one_hot, categorical_features)], remainder="passthrough" ) # 用训练集拟合编码器并转换特征 X_train_transformed = transformer.fit_transform(X_train) # 训练模型 model = RandomForestClassifier() model.fit(X_train_transformed, y_train) # 保存编码器和模型到本地 dump(transformer, 'preprocessor.joblib') dump(model, 'model.joblib')
2. Flask应用加载预保存的编码器
不要在predict路由里重新创建编码器,而是直接加载训练好的规则,仅调用transform转换用户数据:
from joblib import load import pandas as pd from scipy.sparse import csr_matrix from flask import redirect, url_for, render_template from your_module import PredictForm, login_required # 在Flask app初始化时提前加载(避免每次请求重复加载) model = load('model.joblib') preprocessor = load('preprocessor.joblib') @app.route('/predict', methods=['GET','POST']) @login_required def predict(): form = PredictForm() if form.validate_on_submit(): # 省略用户输入的获取与CP值计算逻辑... # 构造用户输入的DataFrame da = {'DFOS':[Dfos],'CP':[cp],'PASSION':[passion]} x = pd.DataFrame(data=da) # 关键:用预训练的编码器转换数据,而非重新拟合 transformed_x = preprocessor.transform(x) x = csr_matrix(transformed_x) prediction = model.predict(x) if prediction is not None: return redirect(url_for('result', prediction=prediction)) return render_template('predict.html', title='predict', form=form)
3. 额外注意事项
- 确保表单提交的特征名称、数据类型和训练集完全一致,比如
DFOS、CP、PASSION的取值范围要和训练时的类别匹配; - 初始化OneHotEncoder时设置
handle_unknown='ignore',可以兼容训练集未出现的用户输入类别,避免编码报错。
内容的提问来源于stack exchange,提问作者Dancun Gerald
相关产品推荐
相关产品推荐

