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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 05:45:32