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

TorchMetrics多分类语义分割Accuracy计算结果差异疑问

TorchMetrics多分类语义分割准确率计算差异解析

问题重现

预测张量定义

# shape: (1, 3, 2, 2) => (batch_size, classes, height, width)
mask_multiclass_pred = torch.tensor(
    [[
        [
            # predictions for first class per pixel
            [0.85, 0.4],
            [0.4, 0.3],
        ],
        [
            # predictions for second class per pixel
            [0, 0.8],
            [0, 1],
        ],
        [
            # predictions for third class per pixel
            [0.8, 0.6],
            [0.7, 0.3],
        ]
    ]],
    dtype=torch.float32
)

转换为类别索引张量

reduced_pred = torch.argmax(mask_multiclass_pred, dim=1)
reduced_pred = torch.where(torch.amax(mask_multiclass_pred, dim=1) >= 0.5, reduced_pred, -1)

得到结果:

# shape: (1, 2, 2) => (batch_size, height, width)
tensor([[[0, 1],
         [2, 1]]])

真实标签张量定义

# shape: (1, 2, 2) => (batch_size, height, width)
mask_multiclass_gt = torch.tensor(
    [
        [
            # class 0, 1, or 2 per pixel => (2, 2) shape for mask
            [0, 1],
            [0, 2],
        ],
    ],
    dtype=torch.int
)

TorchMetrics计算结果

使用MulticlassAccuracy计算类别准确率:

from torchmetrics.classification import MulticlassAccuracy

seg_acc_cls = MulticlassAccuracy(num_classes=3, top_k=1, average="none", multidim_average="global")
seg_acc_cls(mask_multiclass_pred, mask_multiclass_gt)

输出结果:

# shape (3,) => one accuracy per class (3 classes)
tensor([0.5000, 1.0000, 0.0000])

推导矛盾点

手动推导的预期结果与TorchMetrics输出不符:

  • 类别0:预期准确率0.75,实际输出0.5;
  • 类别1:预期准确率0.75,实际输出1.0;
  • 类别2:预期准确率0.5,实际输出0.0。

差异核心原因

两者对类别准确率的定义完全不同:

  • 手动推导逻辑:将每个类别视为「当前类 vs 其他所有类」的二分类任务,计算该二分类任务的全局准确率,公式为:

    准确率 = (TP + TN) / 总样本数
    

    其中TP是「预测为当前类且真实为当前类」的样本数,TN是「预测不为当前类且真实不为当前类」的样本数。

  • TorchMetrics计算逻辑:当设置average="none"时,MulticlassAccuracy计算的是每个类别的召回率(Recall),仅关注真实属于该类的样本中被正确预测的比例,公式为:

    类别c的准确率 = 预测为c且真实为c的样本数 / 真实为c的样本数
    

    该计算完全忽略真实不属于当前类的样本(TN、FP)。

对应示例的具体计算

  • 类别0:真实属于该类的样本共2个(像素1、3),仅1个被正确预测为0,结果为1/2 = 0.5;
  • 类别1:真实属于该类的样本共1个(像素2),全部被正确预测为1,结果为1/1 = 1.0;
  • 类别2:真实属于该类的样本共1个(像素4),未被正确预测为2,结果为0/1 = 0.0。

此计算逻辑完全匹配TorchMetrics的输出结果。

实现预期结果的方法

若要计算每个类别作为二分类任务的全局准确率,可通过以下方式:

  • 手动统计每个类别的TP、TN、FP、FN,自行计算准确率;
  • 对每个类别单独使用BinaryAccuracy,将多分类标签转换为二分类标签(当前类为1,其他类为0);
  • 使用MulticlassConfusionMatrix生成混淆矩阵后,基于矩阵计算每个类别的二分类准确率。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 05:35:26