如何在PyTorch中为14种疾病类别生成统一多标签混淆矩阵
解决多标签分类14x14统一混淆矩阵生成问题
核心思路
sklearn.metrics.multilabel_confusion_matrix会为每个类别单独生成2x2的混淆矩阵,而你需要的全局14x14矩阵,本质是统计真实类别与预测类别之间的共现次数:矩阵中[i][j]位置的数值,表示真实标签包含类别i且预测标签包含类别j的样本总数。
实现步骤(基于PyTorch)
1. 提取所有测试样本的真实标签与预测标签
从test_loader中批量获取真实标签和模型预测结果,转为二进制0/1矩阵:
import torch import numpy as np # 模型切换到评估模式 model.eval() y_true_list = [] y_pred_list = [] device = torch.device("cuda" if torch.cuda.is_available() else "cpu") with torch.no_grad(): for inputs, labels in test_loader: inputs = inputs.to(device) outputs = model(inputs) # 用sigmoid+阈值(比如0.5)生成二进制预测标签 preds = torch.sigmoid(outputs) > 0.5 # 转numpy并存储 y_true_list.append(labels.cpu().numpy()) y_pred_list.append(preds.cpu().numpy()) # 合并所有样本的标签矩阵,shape为(n_samples, 14) y_true = np.concatenate(y_true_list, axis=0) y_pred = np.concatenate(y_pred_list, axis=0)
2. 生成14x14统一混淆矩阵
利用矩阵点积直接计算类别共现次数,得到14x14的全局矩阵:
# y_true.T 是(14, n_samples),y_pred是(n_samples,14),点积后得到(14,14)的混淆矩阵 global_confusion_matrix = y_true.T @ y_pred
矩阵元素含义:
- 对角线
global_confusion_matrix[i][i]:类别i的真阳性(TP)样本数 - 非对角线
global_confusion_matrix[i][j]:真实含类别i且预测含类别j的样本数
3. 可视化矩阵
用seaborn和matplotlib绘制热力图:
import seaborn as sns import matplotlib.pyplot as plt # 替换成你的14个疾病类别名称 class_names = ["疾病1", "疾病2", "疾病3", "疾病4", "疾病5", "疾病6", "疾病7", "疾病8", "疾病9", "疾病10", "疾病11", "疾病12", "疾病13", "疾病14"] plt.figure(figsize=(12, 10)) sns.heatmap(global_confusion_matrix, annot=True, fmt='d', cmap='Blues', xticklabels=class_names, yticklabels=class_names) plt.xlabel('预测类别') plt.ylabel('真实类别') plt.title('多标签全局混淆矩阵(14x14)') plt.show()
补充说明
如果你需要的是其他统计维度的矩阵(比如真实含i但不含j的样本数),可以基于y_true和y_pred的二进制矩阵进行逻辑运算统计,但上述方法是最符合“x轴y轴均为14个类别”需求的常规实现。
内容的提问来源于stack exchange,提问作者Mukhlis Raza
相关产品推荐
相关产品推荐

