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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.05 17:25:18