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

如何生成二类混淆矩阵并导出为JSON文件?解决Sklearn与Torchmetrics报错问题

解决CNN二类混淆矩阵生成及导出JSON的问题

看起来你在生成二类混淆矩阵时遇到了两个典型的工具库适配问题,我来帮你逐个拆解,结合你的代码给出可直接运行的修改方案:

错误原因分析

1. Sklearn的ValueError: Classification metrics can't handle a mix of binary and continuous targets

你使用的是BCEWithLogitsLoss,模型输出的是连续的logits值,但sklearn的confusion_matrix要求输入必须是离散的类别标签(0或1),不能直接传入概率或logits,这就是报错的核心原因。

2. Torchmetrics的AttributeError: 'list' object has no attribute 'shape'

这个错误说明你大概率是把Python列表直接传给了torchmetrics的ConfusionMetric工具,而它只接受PyTorch Tensor作为输入。另外还要注意输入形状的匹配——如果模型输出是(batch_size, 1)的二维张量,需要压缩成(batch_size,)的一维张量,和标签形状对齐。


完整解决方案代码

首先补全你中断的trainCNN函数,然后添加测试阶段的结果收集、混淆矩阵生成和JSON导出逻辑:

import logging
import os
import pickle
import json
from pathlib import Path
from typing import Any, Callable
import matplotlib.pyplot as plt
import numpy as np
import torch
from sklearn.metrics import ConfusionMatrixDisplay, confusion_matrix
from torch import Tensor, nn
from torch.utils.data import DataLoader, random_split
from torch.utils.data import Dataset as _Dataset
from torchmetrics.classification import ConfusionMatrix
from .ecg import build_model_prototype

""" uv run python -m bach_ben.fibrillation """
def get_dataset_dir() -> str:
    p = os.getenv("DATASET_DIR")
    if p is None:
        raise ValueError("DATASET_DIR environment variable not set")
    logger = logging.getLogger(__name__)
    logger.log(level=logging.INFO, msg=f"using DATASET_DIR {p}")
    return p

def get_model_dir() -> str:
    p = os.getenv("MODEL_DIR")
    if p is None:
        raise ValueError("MODEL_DIR environment variable not set")
    logger = logging.getLogger(__name__)
    logger.log(level=logging.INFO, msg=f"using MODEL_DIR {p}")
    return p

class MitBihAtrialFibrillationDataset(_Dataset[tuple[torch.Tensor, torch.Tensor]]):
    def __init__(
        self,
        dataset_root: Path,
        sample_preprocessing: Callable[[torch.Tensor], torch.Tensor] | None = None,
        dtype_samples: torch.dtype = torch.float32,
        dtype_labels: torch.dtype = torch.float32,
    ) -> None:
        super().__init__()
        self._samples, self._labels = self._load_samples_labels(dataset_root)
        self._samples = self._samples.to(dtype_samples)
        self._labels = self._labels.to(dtype_labels)
        if sample_preprocessing is not None:
            self._samples = sample_preprocessing(self._samples)

    def __len__(self) -> int:
        return len(self._labels)

    def __getitem__(self, index: Any) -> tuple[torch.Tensor, torch.Tensor]:
        return self._samples[index], self._labels[index]

    @staticmethod
    def _load_samples_labels(dataset_root: Path) -> tuple[torch.Tensor, torch.Tensor]:
        af_dir = dataset_root / "atrial_fibrillation"
        sinus_dir = dataset_root / "sinus_rhythm"
        for d in (af_dir, sinus_dir):
            if not d.exists():
                raise ValueError(f"{d} does not exist!")

        def load_samples_labels_for_class(
            class_name: str,
        ) -> tuple[torch.Tensor, torch.Tensor]:
            sample_buffer = []
            label = ""
            for pickle_file in (dataset_root / class_name).glob("*.pickle"):
                with pickle_file.open("rb") as in_file:
                    sample, label = pickle.load(in_file)
                sample_buffer.append(torch.tensor(sample.T).unsqueeze(dim=0))
            samples = torch.cat(sample_buffer)
            n = len(samples)
            labels = torch.zeros(n) if label == "(N" else torch.ones(n)
            return samples, labels

        sinus_samples, sinus_labels = load_samples_labels_for_class("sinus_rhythm")
        af_samples, af_labels = load_samples_labels_for_class("atrial_fibrillation")
        samples = torch.cat([sinus_samples, af_samples]).to(torch.float32)
        labels = torch.cat([sinus_labels, af_labels]).to(torch.int64)
        return samples, labels

def trainCNN(model: torch.nn.Module, epochs, batch_size):
    device = "cuda" if torch.cuda.is_available() else "cpu"
    torch.manual_seed(0)
    np.random.seed(0)
    dataset_base_dir = Path(get_dataset_dir()) / "mit-bih-atrial-fibrillation"
    model_save_dir = Path(get_model_dir())

    def select_first_channel(x: torch.Tensor) -> torch.Tensor:
        return x[:, 0:1, :]

    ds = MitBihAtrialFibrillationDataset(dataset_base_dir, sample_preprocessing=select_first_channel)
    train_ds, val_ds, test_ds = random_split(ds, [0.6, 0.2, 0.2])

    mean = 0
    std = 0
    for x, _ in train_ds:  # 修正:遍历训练集样本计算均值方差
        mean += x[0].mean()
        std += x[0].std()
    mean /= len(train_ds)
    std /= len(train_ds)

    def normalize(x: torch.Tensor) -> torch.Tensor:
        return (x - mean) / std

    train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True)
    val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=True)
    test_loader = DataLoader(test_ds, batch_size=batch_size, shuffle=True)

    optim = torch.optim.Adam(model.parameters(), lr=1e-3)
    loss_fn = nn.BCEWithLogitsLoss()

    train_hist, val_hist = [], []
    prev_va_loss = float('inf')  # 初始化早停的验证损失为无穷大
    for epoch in range(1, epochs + 1):
        model.train()
        tr_loss = 0.0
        for x, y in train_loader:
            x, y = normalize(x.to(device)), y.to(device).float()  # 修正:BCE要求标签为float类型
            optim.zero_grad()
            pred = model(x)
            loss = loss_fn(pred, y.unsqueeze(1))  # 修正:对齐预测与标签形状
            loss.backward()
            optim.step()
            tr_loss += loss.item() * x.size(0)
        tr_loss /= len(train_loader.dataset)  # 修正:按总样本数计算平均损失
        if tr_loss < 0.036:
            print(f"Train loss reached threshold {tr_loss:.4f}, stopping early")
            break
        model.eval()
        va_loss = 0.0
        with torch.no_grad():
            for x, y in val_loader:
                x, y = normalize(x.to(device)), y.to(device).float()
                pred = model(x)
                va_loss += loss_fn(pred, y.unsqueeze(1)).item() * x.size(0)
        va_loss /= len(val_loader.dataset)
        train_hist.append(tr_loss)
        val_hist.append(va_loss)

        print(f"Epoch {epoch:2d} | Train Loss: {tr_loss:.4f} | Val Loss: {va_loss:.4f}")

        # 早停逻辑
        if va_loss >= prev_va_loss:
            print("Validation loss did not improve, stopping early")
            break
        prev_va_loss = va_loss

    # --- 测试阶段:收集真实标签和预测结果 ---
    model.eval()
    y_true = []
    y_pred_logits = []
    with torch.no_grad():
        for x, y in test_loader:
            x = normalize(x.to(device))
            pred = model(x)
            y_true.append(y.cpu())
            y_pred_logits.append(pred.cpu())

    # 合并为完整的Tensor
    y_true = torch.cat(y_true)
    y_pred_logits = torch.cat(y_pred_logits)

    # --- 1. 使用Sklearn生成混淆矩阵并导出JSON ---
    # 将logits转换为类别标签(0/1)
    y_pred_probs = torch.sigmoid(y_pred_logits).squeeze(dim=1)
    y_pred = (y_pred_probs > 0.5).int().numpy()
    y_true_np = y_true.numpy()

    cm_sklearn = confusion_matrix(y_true_np, y_pred)
    print("\nSklearn Confusion Matrix:")
    print(cm_sklearn)

    # 导出为JSON
    with open("sklearn_confusion_matrix.json", "w") as f:
        json.dump(cm_sklearn.tolist(), f, indent=4)
    print("Sklearn混淆矩阵已导出为sklearn_confusion_matrix.json")

    # --- 2. 使用Torchmetrics生成混淆矩阵并导出JSON ---
    # 初始化二类混淆矩阵指标
    confmat_torch = ConfusionMatrix(task="binary", num_classes=2)
    # 直接传入logits,自动处理类别转换
    confmat_torch.update(y_pred_logits.squeeze(dim=1), y_true)
    cm_torch = confmat_torch.compute()

    print("\nTorchmetrics Confusion Matrix:")
    print(cm_torch)

    # 导出为JSON
    with open("torchmetrics_confusion_matrix.json", "w") as f:
        json.dump(cm_torch.tolist(), f, indent=4)
    print("Torchmetrics混淆矩阵已导出为torchmetrics_confusion_matrix.json")

    return train_hist, val_hist

关键修正点说明

  1. 数据加载与损失匹配:

    • 修正了BCEWithLogitsLoss的输入形状:标签需要从(batch_size,)转换为(batch_size,1),和模型输出形状对齐
    • 修正了损失计算的平均方式:应该除以数据集总样本数,而不是DataLoader的batch数量
  2. Sklearn混淆矩阵适配:

    • 用torch.sigmoid()将logits转换为概率,再用0.5阈值得到类别标签
    • 将Tensor转换为numpy数组后传入confusion_matrix
  3. Torchmetrics混淆矩阵适配:

    • 确保输入是Tensor而非列表,通过torch.cat()合并批次结果
    • 直接传入logits,工具会自动处理类别转换(因为任务设置为binary)
    • 压缩预测张量的维度,确保和标签形状一致
  4. JSON导出:

    • 用.tolist()将numpy数组/Tensor转换为Python列表,因为JSON不支持直接序列化这些类型

内容的提问来源于stack exchange,提问作者traq

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.28 06:40:22