如何在Keras中为3D脑肿瘤分割模型生成混淆矩阵及ROC曲线?
3D脑肿瘤分割:生成混淆矩阵与ROC曲线解决方案
核心思路
混淆矩阵和ROC曲线属于全局评估指标,无需在model.compile()的metrics参数中定义,只需在模型预测完成后,单独提取测试集的真实标签与预测结果,转换为NumPy数组后用scikit-learn工具生成即可。
步骤1:提取测试集的真实标签与预测结果
先从test_generator中获取所有真实标签,同时用模型生成所有预测结果:
import numpy as np from sklearn.metrics import confusion_matrix, roc_curve, auc from sklearn.preprocessing import label_binarize # 获取所有测试集的真实标签(假设y是one-hot编码,转成类别索引) y_true = [] for x, y in test_generator: y_true.extend(np.argmax(y, axis=-1).flatten()) y_true = np.array(y_true) # 生成预测概率,再转换为类别索引 y_pred_proba = model.predict(test_generator, verbose=1) y_pred = np.argmax(y_pred_proba, axis=-1).flatten()
若
test_generator输出的y已是类别索引而非one-hot编码,去掉np.argmax(y, axis=-1)即可。
步骤2:生成混淆矩阵
用scikit-learn的confusion_matrix直接计算,还可可视化:
# 计算混淆矩阵 cm = confusion_matrix(y_true, y_pred) print("混淆矩阵:") print(cm) # 可视化混淆矩阵 import matplotlib.pyplot as plt import seaborn as sns plt.figure(figsize=(8,6)) sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=['背景', '坏死区', '水肿区', '增强区'], yticklabels=['背景', '坏死区', '水肿区', '增强区']) plt.xlabel('预测类别') plt.ylabel('真实类别') plt.title('3D脑肿瘤分割混淆矩阵') plt.show()
步骤3:生成多分类ROC曲线
针对4分类任务,采用one-vs-rest策略生成ROC曲线:
# 将真实标签二值化(one-vs-rest格式) y_true_bin = label_binarize(y_true, classes=[0,1,2,3]) n_classes = y_true_bin.shape[1] # 计算每个类别的ROC曲线与AUC值 fpr = dict() tpr = dict() roc_auc = dict() for i in range(n_classes): fpr[i], tpr[i], _ = roc_curve(y_true_bin[:, i], y_pred_proba[:, i]) roc_auc[i] = auc(fpr[i], tpr[i]) # 绘制ROC曲线 plt.figure(figsize=(10,8)) colors = ['blue', 'red', 'green', 'orange'] class_names = ['背景', '坏死区', '水肿区', '增强区'] for i, color in zip(range(n_classes), colors): plt.plot(fpr[i], tpr[i], color=color, lw=2, label=f'ROC曲线 (类别 {class_names[i]}, AUC = {roc_auc[i]:.2f})') plt.plot([0, 1], [0, 1], 'k--', lw=2) plt.xlim([0.0, 1.0]) plt.ylim([0.0, 1.05]) plt.xlabel('假阳性率') plt.ylabel('真阳性率') plt.title('多分类ROC曲线') plt.legend(loc="lower right") plt.show()
常见问题解决
- Tensor转NumPy失败:TensorFlow 2.x默认开启Eager模式,直接调用
.numpy()即可转换,GPU上的Tensor会自动同步到CPU:# 示例:自定义函数中转换Tensor def custom_func(y_true, y_pred): y_true_np = y_true.numpy() y_pred_np = y_pred.numpy() # 后续处理逻辑 - TypeError: 'property' object is not iterable:通常是误将指标属性而非实例传入
metrics参数,比如应写tf.keras.metrics.Precision()而非precision。
内容的提问来源于stack exchange,提问作者Rayhan Papar
相关产品推荐
相关产品推荐

