如何生成二分类混淆矩阵并导出为JSON?解决Sklearn/Torchmetrics报错
解决二分类混淆矩阵生成及JSON导出问题
错误原因分析
Sklearn报错:ValueError: Classification metrics can't handle a mix of binary and continuous targets
- 模型使用
BCEWithLogitsLoss,输出是未经过sigmoid的连续logits值,但sklearn的confusion_matrix需要离散的类别标签(0/1),而非连续概率值。 y_true和y_pred是存储张量的列表,未拼接成一维数组,sklearn无法直接处理。
Torchmetrics报错:AttributeError: 'list' object has no attribute 'shape'
- Torchmetrics的
ConfusionMatrix要求输入张量,但传入的是存储张量的列表,未拼接为完整张量,因此找不到shape属性。
解决步骤
1. 处理标签与预测结果
- 对模型输出的logits应用
sigmoid,转换为0-1区间的概率值,再通过阈值(如0.5)得到类别标签。 - 将
y_true和y_pred列表中的张量拼接为一维张量,再转为numpy数组或保持张量格式。
2. 生成混淆矩阵(两种方式任选)
3. 将混淆矩阵导出为JSON文件
修改后的测试部分代码
# 替换原代码中测试环节的代码 test_loss = 0.0 y_true, y_pred_logits = [], [] model.eval() with torch.no_grad(): for x, y in test_loader: x, y = normalize(x.to(device)), y.to(device) output = model(x) test_loss += loss_fn(output, y).item() * x.size(0) # 收集真实标签和模型输出(转到CPU避免设备不匹配) y_true.append(y.cpu()) y_pred_logits.append(output.cpu()) # 拼接张量并处理为可计算混淆矩阵的格式 y_true = torch.cat(y_true).numpy() # 将logits转为概率值,再通过阈值得到类别标签 y_pred_probs = torch.sigmoid(torch.cat(y_pred_logits)).numpy() y_pred = (y_pred_probs >= 0.5).astype(int) # -------------------------- 用Sklearn生成混淆矩阵 -------------------------- from sklearn.metrics import confusion_matrix cm = confusion_matrix(y_true, y_pred) print("Sklearn混淆矩阵:") print(cm) # -------------------------- 用Torchmetrics生成混淆矩阵 -------------------------- from torchmetrics.classification import ConfusionMatrix cm_metric = ConfusionMatrix(task="binary", num_classes=2) cm_torch = cm_metric(torch.tensor(y_pred), torch.tensor(y_true)) print("\nTorchmetrics混淆矩阵:") print(cm_torch.numpy()) # -------------------------- 导出为JSON文件 -------------------------- import json # 将numpy数组转为列表以支持JSON序列化 cm_data = { "confusion_matrix": cm.tolist(), "class_labels": ["sinus_rhythm", "atrial_fibrillation"] } with open("confusion_matrix.json", "w") as f: json.dump(cm_data, f, indent=4) print("\n混淆矩阵已导出到confusion_matrix.json")
注意事项
- 阈值0.5是二分类默认值,可根据模型ROC曲线调整最优阈值。
- 如果模型最后一层已添加
sigmoid激活,无需再调用torch.sigmoid(),直接用输出做概率判断即可。 - 确保真实标签
y_true为整数类型(你的数据集加载时已转为torch.int64,转numpy后符合要求)。
内容的提问来源于stack exchange,提问作者traq
相关产品推荐
相关产品推荐

