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)
问题根源
原代码的硬伤都是没考虑多分类的特殊性:
- 直接取
predict_proba[:,1]只适用于二分类,多分类下需要获取所有类别的概率 - 敏感度/特异度的二分类公式直接套用到多分类会报错,因为混淆矩阵维度发生变化
- ROC曲线和AUC计算没有适配多分类的One-vs-Rest逻辑
- 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))
关键修改说明
- 数据集替换:用鸢尾花(3类别)测试多分类逻辑,原代码的心脏病数据集被处理成二分类,无法验证多分类效果
- SVM参数调整:添加
probability=True让SVM输出概率,便于AUC计算 - 指标适配多分类:
- 敏感度/特异度:遍历每个类别,按One-vs-Rest逻辑计算后取宏观平均
- 平衡准确率:直接使用sklearn内置的
balanced_accuracy_score,比手动计算更可靠 - AUC:多分类下启用
multi_class='ovr'参数,支持宏观/微观两种平均方式
- ROC曲线优化:多分类下为每个类别单独绘制ROC曲线,清晰展示每个类别的区分能力
- 标签处理:用
LabelBinarizer将多分类标签转为二进制格式,适配ROC计算逻辑
内容的提问来源于stack exchange,提问作者Daniel Shiverman
相关产品推荐
相关产品推荐

