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
相关产品推荐
相关产品推荐

