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

如何修改Pickle加载的XGBoost模型默认分类阈值至0.7?

如何修改XGBoost模型默认分类阈值为0.7并解决super()调用错误

你遇到的RuntimeError: super(): __class__ cell not found,是因为动态绑定的独立函数没有类定义时生成的__class__闭包上下文,导致super()无法正确定位父类。以下是两种可行的解决方案:

方案1:修正自定义predict方法,避免super()调用

首先纠正你原方法的逻辑错误:super().predict()返回的是类别标签而非概率值,应直接调用predict_proba()获取概率后应用阈值。同时去掉super()调用,彻底规避上下文问题:

import numpy as np
import types
import pickle

def new_predict(
    self,
    X,
    output_margin=False,
    ntree_limit=None,
    validate_features=True,
    base_margin=None,
    iteration_range=None,
) -> np.ndarray:
    if output_margin:
        # 输出原始边际值,直接调用原模型方法
        return self.predict(X, output_margin=True, ntree_limit=ntree_limit, 
                           validate_features=validate_features, base_margin=base_margin,
                           iteration_range=iteration_range)
    
    # 获取二分类概率值
    class_probs = self.predict_proba(
        X=X,
        ntree_limit=ntree_limit,
        validate_features=validate_features,
        base_margin=base_margin,
        iteration_range=iteration_range,
    )

    if self.n_classes_ != 2:
        # 多分类场景保持原逻辑
        column_indexes = np.argmax(class_probs, axis=1)
    else:
        # 二分类应用0.7阈值
        column_indexes = np.zeros(class_probs.shape[0], dtype=int)
        column_indexes[class_probs[:, 1] >= 0.7] = 1

    if hasattr(self, "_le"):
        return self._le.inverse_transform(column_indexes)
    return column_indexes

# 加载模型并绑定自定义方法
with open("XGBoost_model.pkl", "rb") as fr:
    model = pickle.load(fr)

model.predict = types.MethodType(new_predict, model)

# 验证结果一致性
Y_pred_2 = model.predict(X_test)
Y_pred_1 = (model.predict_proba(X_test)[:, 1] >= 0.7).astype(int)
assert np.array_equal(Y_pred_1, Y_pred_2)

方案2:通过子类继承实现(更规范)

提前定义继承自XGBClassifier的子类,重写predict方法,再将加载后的模型实例的类替换为该子类:

from xgboost import XGBClassifier
import numpy as np
import pickle

class XGBClassifierWithThreshold(XGBClassifier):
    def __init__(self, threshold=0.7, **kwargs):
        super().__init__(**kwargs)
        self.threshold = threshold
    
    def predict(
        self,
        X,
        output_margin=False,
        ntree_limit=None,
        validate_features=True,
        base_margin=None,
        iteration_range=None,
    ) -> np.ndarray:
        if output_margin:
            return super().predict(X, output_margin=True, ntree_limit=ntree_limit,
                                  validate_features=validate_features, base_margin=base_margin,
                                  iteration_range=iteration_range)
        
        class_probs = super().predict_proba(X, ntree_limit=ntree_limit,
                                           validate_features=validate_features, base_margin=base_margin,
                                           iteration_range=iteration_range)
        
        if self.n_classes_ != 2:
            column_indexes = np.argmax(class_probs, axis=1)
        else:
            column_indexes = np.zeros(class_probs.shape[0], dtype=int)
            column_indexes[class_probs[:, 1] >= self.threshold] = 1
        
        if hasattr(self, "_le"):
            return self._le.inverse_transform(column_indexes)
        return column_indexes

# 加载模型并替换类
with open("XGBoost_model.pkl", "rb") as fr:
    model = pickle.load(fr)

model.__class__ = XGBClassifierWithThreshold
model.threshold = 0.7

# 验证结果一致性
Y_pred_2 = model.predict(X_test)
Y_pred_1 = (model.predict_proba(X_test)[:, 1] >= 0.7).astype(int)
assert np.array_equal(Y_pred_1, Y_pred_2)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 02:20:37