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

Python逻辑回归二分类模型类继承封装后输出异常求解

逻辑回归类重构问题修复方案

核心错误原因

  • 继承初始化错误:Logit子类的__init__方法调用父类构造函数的语法错误,且未传入父类要求的scaled_dataset、label、new_variable_name三个必填参数
  • 变量作用域错误:split方法中定义的X_train、X_val、y_train、y_val是局部变量,其他类方法无法访问,需要改为实例属性(加self.前缀)
  • 方法未主动调用:主程序仅实例化了Logit类并打印对象,没有执行数据集拆分、模型训练、预测、绘图等业务方法,自然不会输出预期结果
  • 方法定义不兼容:父类Model中LR_model是普通方法,子类中重写为property属性,父类初始化时调用self.LR_model()的逻辑会报错
  • 之前输出的<bound method xxx>类型内容,是因为直接打印方法对象没有加()执行导致的,正确调用方法即可消除

修复后完整代码

import os
import pandas as pd
import numpy as np
import matplotlib.pyplot as plt
import seaborn as sns
from sklearn.model_selection import train_test_split, GridSearchCV
from sklearn.linear_model import LogisticRegression
from sklearn import metrics
from sklearn.metrics import roc_curve, plot_roc_curve

# 全局数据配置
datafile = pd.read_csv(r'diabetes_dataset.csv')  
label = datafile['Outcome']
cols = list(datafile.columns[:-1])   
variable_name = 'Outcome'  
main_dir = 'Final_folder'  
output_folder = os.path.join(main_dir, 'output')
os.makedirs(output_folder, exist_ok=True)

scaled_dataset = pd.read_csv('scaled_dataset.csv')
new_variable_name = 'Label'


class Model:
    def __init__(self, scaled_dataset, label, new_variable_name):
        self.df = scaled_dataset
        self.label = label
        self.new_variable_name = new_variable_name
        # 提前声明实例属性
        self.X_train = None
        self.X_val = None
        self.y_train = None
        self.y_val = None
        self.model_name = self.LR_model()

    def split(self):
        y = self.label.values
        X = self.df.drop(labels=self.new_variable_name, axis=1).values
        self.X_train, self.X_val, self.y_train, self.y_val = train_test_split(X, y, test_size=0.2, shuffle=True, random_state=42)  
        print("训练集样本数:", self.y_train.shape[0], "- 测试集样本数:", self.y_val.shape[0])

    def fitting(self):
        return self.model_name.fit(self.X_train, self.y_train)

    def train_predict(self):
        predicted = self.model_name.predict(self.X_train)
        print("训练集分类报告 \n %s:\n\n%s\n"
              % (self.model_name, metrics.classification_report(self.y_train, predicted)))
        print("训练集AUC:\n%s" % metrics.roc_auc_score(self.y_train,
               self.model_name.predict_proba(self.X_train)[:, 1]))

        class_names = [0, 1]
        fig, ax = plt.subplots()
        tick_marks = np.arange(len(class_names))
        plt.xticks(tick_marks, class_names)
        plt.yticks(tick_marks, class_names)
        sns.heatmap(pd.DataFrame(metrics.confusion_matrix(self.y_train, predicted)), annot=True, cmap="YlGnBu", fmt='g')
        ax.xaxis.set_label_position("top")
        plt.tight_layout()
        plt.title('训练集混淆矩阵', y=1.1)
        plt.ylabel('真实标签')
        plt.xlabel('预测标签')
        plt.savefig(os.path.join(output_folder, 'train_confusion_matrix.png'))
        plt.show()

    def val_predict(self):
        predicted_val = self.model_name.predict(self.X_val)
        print("\n验证集分类报告 \n %s:\n\n%s\n"
              % (self.model_name, metrics.classification_report(self.y_val, predicted_val)))
        print("验证集AUC:\n%s" % metrics.roc_auc_score(self.y_val, self.model_name.predict_proba(self.X_val)[:, 1]))

        class_names = [0, 1]
        fig, ax = plt.subplots()
        tick_marks = np.arange(len(class_names))
        plt.xticks(tick_marks, class_names)
        plt.yticks(tick_marks, class_names)
        sns.heatmap(pd.DataFrame(metrics.confusion_matrix(self.y_val, predicted_val)), annot=True, cmap="YlGnBu", fmt='g')
        ax.xaxis.set_label_position("top")
        plt.tight_layout()
        plt.title('验证集混淆矩阵', y=1.1)
        plt.ylabel('真实标签')
        plt.xlabel('预测标签')
        plt.savefig(os.path.join(output_folder, 'val_confusion_matrix.png'))
        plt.show()

    def val_roc_curve(self):
        prob_test = self.model_name.predict_proba(self.X_val)
        fpr, tpr, thresholds = roc_curve(self.y_val, prob_test[:, 1])
        plot_roc_curve(self.model_name, self.X_val, self.y_val)
        plt.title('验证集ROC曲线')
        plt.savefig(os.path.join(output_folder, 'val_roc_curve.png'))
        plt.show()

    def LR_model(self):
        pass


class Logit(Model):
    def __init__(self, scaled_dataset, label, new_variable_name):
        # 正确调用父类构造函数传参
        super().__init__(scaled_dataset, label, new_variable_name)
        self.classifier = LogisticRegression(max_iter=1000)
        self.parameters = {'C': [1e-4, 1e-3, 1e-2, 1e-1, 1, 10],
                           'penalty': ['l1', 'l2', 'elasticnet', 'none'],
                           'solver': ['newton-cg', 'lbfgs', 'liblinear', 'sag', 'saga']}

    # 改为普通方法和父类兼容
    def LR_model(self):
        # 先拆分数据集再训练
        self.split()
        CV_modelLR = GridSearchCV(estimator=self.classifier,
                                  param_grid=self.parameters,
                                  cv=3, verbose=2)
        CV_modelLR.fit(self.X_train, self.y_train)
        best_params = CV_modelLR.best_params_
        print(f"最优超参数:{best_params}")
        logit = LogisticRegression(penalty=best_params['penalty'],
                                   C=best_params['C'],
                                   solver=best_params['solver'],
                                   class_weight='balanced',
                                   max_iter=1000)
        logit.fit(self.X_train, self.y_train)
        return logit


if __name__ == '__main__':
    # 实例化时传入父类需要的三个参数
    model = Logit(scaled_dataset, label, new_variable_name)
    # 调用需要的业务方法
    model.train_predict()
    model.val_predict()
    model.val_roc_curve()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.27 23:57:00