如何修改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
相关产品推荐
相关产品推荐

