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

Scikit-learn二分类模型转3类多分类的适配问题求助

多分类任务下Scikit-learn代码适配问题解决

原代码在二分类场景运行正常,但碰到3类以上的多分类任务就彻底失效。需要计算的指标包括准确率、敏感度、特异度、平衡准确率、精确率、召回率、F1值、AUC,还要绘制ROC曲线。试过扁平化变量、循环执行一对一二分类取平均、拆分数据集两两分类这些方法,都没解决问题。

原代码如下:

# data import
from ucimlrepo import fetch_ucirepo
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
from sklearn.model_selection import train_test_split
from sklearn.tree import DecisionTreeClassifier
from sklearn.naive_bayes import GaussianNB
from sklearn.ensemble import RandomForestClassifier
from sklearn.svm import SVC
from sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score, confusion_matrix
from sklearn.metrics import ConfusionMatrixDisplay
from sklearn.metrics import roc_curve, roc_auc_score
from sklearn.preprocessing import LabelBinarizer

# fetch breast cancer dataset
bc = fetch_ucirepo(id=17)

# data (as pandas dataframes)
bc_X = bc.data.features
bc_y = bc.data.targets

# fetch heart disease dataset
hd = fetch_ucirepo(id=45)
# data (as pandas dataframes)
hd_X = hd.data.features
hd_y = hd.data.targets


# fetch iris dataset
ir = fetch_ucirepo(id=53)
# data (as pandas dataframes)
ir_X = ir.data.features
ir_y = ir.data.targets


# fetch wine quality dataset
wq = fetch_ucirepo(id=186)
# data (as pandas dataframes)
wq_X = wq.data.features
wq_y = wq.data.targets


# Step 2: Split the data into training and testing sets
X_train, X_test, y_train, y_test = train_test_split(hd_X, hd_y, test_size=0.2, random_state=42)

# cap values function
def cap_values(values):
    # Convert input to a numpy array
    values = np.array(values)
    # Cap values at 1 using numpy's clip function
    capped_values = np.clip(values, None, 1)
    return capped_values

# Find the indices of rows with missing values in X_train
missing_indices = X_train.isnull().any(axis=1)
missing_indices2 = X_test.isnull().any(axis=1)

# Remove rows with missing values from X_train and y_train
X_train = X_train[~missing_indices]
y_train = y_train[~missing_indices]
X_test = X_test[~missing_indices2]
y_test = y_test[~missing_indices2]

# Apply the cap_values function to y_test and y_train
y_test = cap_values(y_test)
y_train = cap_values(y_train)

# Step 3: Fit each model to the training data

# Decision Tree
dt_model = DecisionTreeClassifier(random_state=42)
dt_model.fit(X_train, y_train)

# Naïve Bayes
nb_model = GaussianNB()
nb_model.fit(X_train, y_train)

# Random Forest
rf_model = RandomForestClassifier(random_state=42)
rf_model.fit(X_train, y_train)

# Support Vector Machine
svm_model = SVC(random_state=42)
svm_model.fit(X_train, y_train)

# Step 4: Evaluate the models using the testing data
models = {
    'Decision Tree': dt_model,
    'Naïve Bayes': nb_model,
    'Random Forest': rf_model,
    'Support Vector Machine': svm_model
}

# Initialize a dictionary to store the performance metrics
performance_metrics = {}

# Iterate through each model and calculate the metrics
for name, model in models.items():
    # Get predictions
    y_pred = model.predict(X_test)
    y_pred_prob = model.predict_proba(X_test)[:, 1] if hasattr(model, 'predict_proba') else model.decision_function(X_test)

    # Confusion matrix
    cm = confusion_matrix(y_test, y_pred)
    ConfusionMatrixDisplay(cm, display_labels=model.classes_).plot(cmap=plt.cm.Blues)
    plt.title(f'Confusion Matrix: {name}')
    plt.show()

    # Calculate sensitivity (recall) and specificity
    if cm.shape == (2, 2):
      # Binary classification metrics
      tn, fp, fn, tp = cm.ravel()
      sensitivity = tp / (tp + fn)
      specificity = tn / (tn + fp)
    else:
      # Multi-class classification metrics
      sensitivity = recall_score(y_test, y_pred, average='macro')
      specificity = None
    sensitivity = tp / (tp + fn)
    specificity = tn / (tn + fp)

    # Calculate balanced accuracy
    balanced_accuracy = (sensitivity + specificity) / 2

    # Calculate precision, recall, and F1-score
    precision = precision_score(y_test, y_pred, average='macro')
    recall = recall_score(y_test, y_pred, average='macro')
    f1 = f1_score(y_test, y_pred, average='macro')

    # Calculate AUC
    fpr, tpr, thresholds = roc_curve(y_test, y_pred_prob)
    auc = roc_auc_score(y_test, y_pred_prob)

    # Store the performance metrics
    performance_metrics[name] = {
        'Accuracy': accuracy_score(y_test, y_pred),
        'Sensitivity': sensitivity,
        'Specificity': specificity,
        'Balanced Accuracy': balanced_accuracy,
        'Precision': precision,
        'Recall': recall,
        'F1 Score': f1,
        'AUC': auc
    }

    # Plot ROC curve
    plt.figure()
    plt.plot(fpr, tpr, label=f'{name} (AUC = {auc:.2f})')
    plt.plot([0, 1], [0, 1], 'k--')  # Diagonal line
    plt.title(f'ROC Curve: {name}')
    plt.xlabel('False Positive Rate')
    plt.ylabel('True Positive Rate')
    plt.legend()
    plt.show()

# Display the performance metrics in a DataFrame
metrics_df = pd.DataFrame(performance_metrics).transpose()
print(metrics_df)

问题根源

原代码的硬伤都是没考虑多分类的特殊性:

  1. 直接取predict_proba[:,1]只适用于二分类,多分类下需要获取所有类别的概率
  2. 敏感度/特异度的二分类公式直接套用到多分类会报错,因为混淆矩阵维度发生变化
  3. ROC曲线和AUC计算没有适配多分类的One-vs-Rest逻辑
  4. SVM默认不输出概率,需要手动开启对应参数

修复后的完整代码

from ucimlrepo import fetch_ucirepo
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
from sklearn.model_selection import train_test_split
from sklearn.tree import DecisionTreeClassifier
from sklearn.naive_bayes import GaussianNB
from sklearn.ensemble import RandomForestClassifier
from sklearn.svm import SVC
from sklearn.metrics import (accuracy_score, precision_score, recall_score, 
                             f1_score, confusion_matrix, balanced_accuracy_score,
                             roc_curve, roc_auc_score)
from sklearn.metrics import ConfusionMatrixDisplay
from sklearn.preprocessing import LabelBinarizer
from sklearn.utils.multiclass import unique_labels

# 加载多分类数据集(鸢尾花,3类别)
ir = fetch_ucirepo(id=53)
X = ir.data.features
y = ir.data.targets.values.ravel()  # 转成一维数组避免维度问题

# 拆分数据集
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)

# 处理缺失值(鸢尾花无缺失,保留通用逻辑)
missing_indices = X_train.isnull().any(axis=1)
X_train = X_train[~missing_indices]
y_train = y_train[~missing_indices]

missing_indices_test = X_test.isnull().any(axis=1)
X_test = X_test[~missing_indices_test]
y_test = y_test[~missing_indices_test]

# 初始化模型,SVM需开启probability=True以输出概率
models = {
    'Decision Tree': DecisionTreeClassifier(random_state=42),
    'Naïve Bayes': GaussianNB(),
    'Random Forest': RandomForestClassifier(random_state=42),
    'Support Vector Machine': SVC(probability=True, random_state=42)
}

# 存储指标
performance_metrics = {}

# 标签二值化,适配多分类ROC计算
lb = LabelBinarizer()
y_test_bin = lb.fit_transform(y_test)
n_classes = len(lb.classes_)

for name, model in models.items():
    # 训练模型
    model.fit(X_train, y_train)
    
    # 获取预测结果和概率
    y_pred = model.predict(X_test)
    if hasattr(model, 'predict_proba'):
        y_pred_prob = model.predict_proba(X_test)
    else:
        # 无predict_proba时,用decision_function并归一化
        y_pred_prob = model.decision_function(X_test)
        y_pred_prob = np.exp(y_pred_prob) / np.sum(np.exp(y_pred_prob), axis=1, keepdims=True)
    
    # 绘制混淆矩阵
    cm = confusion_matrix(y_test, y_pred)
    disp = ConfusionMatrixDisplay(confusion_matrix=cm, display_labels=lb.classes_)
    disp.plot(cmap=plt.cm.Blues)
    plt.title(f'混淆矩阵: {name}')
    plt.show()
    
    # 计算通用指标
    accuracy = accuracy_score(y_test, y_pred)
    balanced_acc = balanced_accuracy_score(y_test, y_pred)
    precision_macro = precision_score(y_test, y_pred, average='macro')
    recall_macro = recall_score(y_test, y_pred, average='macro')
    f1_macro = f1_score(y_test, y_pred, average='macro')
    
    # 按One-vs-Rest逻辑计算每个类别的敏感度和特异度,再取宏观平均
    sensitivity_list = []
    specificity_list = []
    classes = unique_labels(y_test, y_pred)
    for cls in classes:
        tp = np.sum((y_test == cls) & (y_pred == cls))
        fn = np.sum((y_test == cls) & (y_pred != cls))
        fp = np.sum((y_test != cls) & (y_pred == cls))
        tn = np.sum((y_test != cls) & (y_pred != cls))
        
        sensitivity = tp / (tp + fn) if (tp + fn) != 0 else 0
        specificity = tn / (tn + fp) if (tn + fp) != 0 else 0
        
        sensitivity_list.append(sensitivity)
        specificity_list.append(specificity)
    
    sensitivity_macro = np.mean(sensitivity_list)
    specificity_macro = np.mean(specificity_list)
    
    # 计算多分类AUC(One-vs-Rest)
    if n_classes == 2:
        auc = roc_auc_score(y_test, y_pred_prob[:, 1])
    else:
        auc_macro = roc_auc_score(y_test, y_pred_prob, multi_class='ovr', average='macro')
        auc_micro = roc_auc_score(y_test, y_pred_prob, multi_class='ovr', average='micro')
    
    # 存储指标
    metrics = {
        '准确率': accuracy,
        '敏感度(宏观平均)': sensitivity_macro,
        '特异度(宏观平均)': specificity_macro,
        '平衡准确率': balanced_acc,
        '精确率(宏观平均)': precision_macro,
        '召回率(宏观平均)': recall_macro,
        'F1值(宏观平均)': f1_macro
    }
    if n_classes > 2:
        metrics['AUC(宏观平均)'] = auc_macro
        metrics['AUC(微观平均)'] = auc_micro
    else:
        metrics['AUC'] = auc
    
    performance_metrics[name] = metrics
    
    # 绘制多分类ROC曲线(One-vs-Rest)
    plt.figure(figsize=(8, 6))
    if n_classes == 2:
        fpr, tpr, _ = roc_curve(y_test, y_pred_prob[:, 1])
        plt.plot(fpr, tpr, label=f'{name} (AUC = {auc:.2f})')
    else:
        for i in range(n_classes):
            fpr, tpr, _ = roc_curve(y_test_bin[:, i], y_pred_prob[:, i])
            plt.plot(fpr, tpr, label=f'类别 {lb.classes_[i]} (AUC = {roc_auc_score(y_test_bin[:, i], y_pred_prob[:, i]):.2f})')
    
    plt.plot([0, 1], [0, 1], 'k--', label='随机猜测')
    plt.title(f'ROC曲线: {name}')
    plt.xlabel('假阳性率')
    plt.ylabel('真阳性率')
    plt.legend()
    plt.show()

# 输出指标表格
metrics_df = pd.DataFrame(performance_metrics).transpose()
print(metrics_df.round(4))

关键修改说明

  1. 数据集替换:用鸢尾花(3类别)测试多分类逻辑,原代码的心脏病数据集被处理成二分类,无法验证多分类效果
  2. SVM参数调整:添加probability=True让SVM输出概率,便于AUC计算
  3. 指标适配多分类:
    • 敏感度/特异度:遍历每个类别,按One-vs-Rest逻辑计算后取宏观平均
    • 平衡准确率:直接使用sklearn内置的balanced_accuracy_score,比手动计算更可靠
    • AUC:多分类下启用multi_class='ovr'参数,支持宏观/微观两种平均方式
  4. ROC曲线优化:多分类下为每个类别单独绘制ROC曲线,清晰展示每个类别的区分能力
  5. 标签处理:用LabelBinarizer将多分类标签转为二进制格式,适配ROC计算逻辑

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 20:58:09