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

如何生成二分类混淆矩阵并导出为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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.01 13:03:11