如何基于DenseNet121模型计算ROC、灵敏度与特异度
乳腺癌图像分类:ROC曲线与灵敏度计算问题
我用DenseNet121进行乳腺癌图像分类任务,已经通过confusion_matrix、classification_report和accuracy_score完成基础评估,但始终无法计算出ROC曲线、灵敏度等指标,尝试多种方法均未成功。
模型构建与预测代码
import tensorflow as tf from tensorflow.keras.layers import Flatten, Dense from tensorflow.keras.models import Model import numpy as np from sklearn.metrics import confusion_matrix, classification_report, accuracy_score in_model = tf.keras.applications.DenseNet121(input_shape=(224,224,3), include_top=False, weights='imagenet',classes = 2) in_model.trainable = False inputs = tf.keras.Input(shape=(224,224,3)) x = in_model(inputs) flat = Flatten()(x) dense_1 = Dense(4096,activation = 'relu')(flat) dense_2 = Dense(4096,activation = 'relu')(dense_1) prediction = Dense(2,activation = 'softmax')(dense_2) in_pred = Model(inputs = inputs,outputs = prediction) in_pred.evaluate(test_data,test_labels) test_ = in_pred.predict(test_data) # 注:原代码中test_text应为test_data y_true = np.argmax(test_labels, axis=1) # 修正命名:真实类别标签 y_pred_class = np.argmax(test_, axis=1) # 修正命名:预测类别标签
已使用的基础评估代码
print(confusion_matrix(y_true, y_pred_class)) print(classification_report(y_true, y_pred_class)) print(accuracy_score(y_true, y_pred_class))
解决方法
1. 灵敏度(召回率)计算
灵敏度本质就是二分类任务中的召回率,你可以直接从classification_report输出的recall字段获取。如果需要单独计算,可使用sklearn的recall_score:
from sklearn.metrics import recall_score # 二分类场景,指定正类别(例:类别1为阳性样本) sensitivity = recall_score(y_true, y_pred_class, pos_label=1) print(f"灵敏度(召回率):{sensitivity:.4f}")
也可通过混淆矩阵手动计算:
cm = confusion_matrix(y_true, y_pred_class) TN, FP, FN, TP = cm.ravel() sensitivity = TP / (TP + FN) print(f"灵敏度:{sensitivity:.4f}")
2. ROC曲线与AUC计算
ROC曲线需要基于预测概率值而非类别标签,所以要直接使用softmax输出的概率(不要用np.argmax转换后的类别):
from sklearn.metrics import roc_curve, roc_auc_score import matplotlib.pyplot as plt # 提取正类(例:类别1)的预测概率 y_pred_proba = test_[:, 1] # 计算ROC曲线的假阳性率(FPR)、真阳性率(TPR)和阈值 fpr, tpr, thresholds = roc_curve(y_true, y_pred_proba) # 计算AUC值 auc_score = roc_auc_score(y_true, y_pred_proba) # 绘制ROC曲线 plt.plot(fpr, tpr, label=f'AUC = {auc_score:.4f}') plt.plot([0, 1], [0, 1], 'k--', label='随机猜测') plt.xlabel('假阳性率(FPR)') plt.ylabel('真阳性率(TPR,即灵敏度)') plt.title('ROC曲线') plt.legend() plt.show() print(f"AUC值:{auc_score:.4f}")
注意事项
- 确认
test_labels格式:如果test_labels已经是类别索引(非one-hot编码),无需np.argmax(test_labels, axis=1),直接用y_true = test_labels即可。 - 原代码中
test_text属于笔误,需修正为test_data。
内容的提问来源于stack exchange,提问作者Eda
相关产品推荐
相关产品推荐

