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

如何正确定义加载预训练随机森林模型的预测类?

问题分析与解决方案

你的错误根源在于:继承RandomForestClassifier后,直接调用父类的predict方法,但当前类实例并未经过拟合(你是加载预训练模型,而非调用fit训练);同时load_model方法无返回值,导致self.model被赋值为None,逻辑完全错位。

下面提供两种可行的修正方案:


方案一:简单封装类(推荐新手使用)

无需继承RandomForestClassifier,直接将预训练模型作为类的属性封装,逻辑更清晰:

import os
import joblib
import numpy as np

class MODEL_RF:
    def __init__(self, model_path):
        # 初始化时直接加载预训练模型
        self.model = joblib.load(os.path.join(model_path, 'rf.pkl'))
    
    def get_pred(self, df):
        validation_features = np.array(df)
        # 调用加载好的模型的预测方法
        pred = self.model.predict(validation_features)
        predict_prob = self.model.predict_proba(validation_features)
        return pred, predict_prob

# 使用示例
model_m = MODEL_RF(r"the path of the model")
prediction, probs = model_m.get_pred(input_df)

方案二:继承RandomForestClassifier(保留sklearn Estimator特性)

如果需要让你的类保留sklearn分类器的原生方法,可以通过复制预训练模型的属性到当前实例:

from sklearn.ensemble import RandomForestClassifier
import os
import joblib
import numpy as np

class MODEL_RF(RandomForestClassifier):
    def __init__(self, model_path, **kwargs):
        # 调用父类构造函数初始化基础参数
        super().__init__(**kwargs)
        # 加载预训练模型
        pretrained_model = joblib.load(os.path.join(model_path, 'rf.pkl'))
        # 将预训练模型的所有属性复制到当前实例
        self.__dict__.update(pretrained_model.__dict__)
    
    def get_pred(self, df):
        validation_features = np.array(df)
        # 直接调用父类的预测方法(此时实例已具备预训练状态)
        pred = self.predict(validation_features)
        predict_prob = self.predict_proba(validation_features)
        return pred, predict_prob

# 使用示例
model_m = MODEL_RF(r"the path of the model")
prediction, probs = model_m.get_pred(input_df)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.24 12:15:36