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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 12:25:14