Torchmetrics多分类任务中F1/Precision/Recall与Accuracy结果一致问题求助
问题原因与解决方法
问题根源
在torchmetrics 0.11.1版本中,F1Score、Precision、Recall这类多分类指标的默认average参数为'micro'。采用micro平均计算时:
- 会将所有类别的真阳性(TP)、假阳性(FP)、假阴性(FN)全局累加后统一计算指标
- 此时Precision、Recall、F1的结果完全相等,并且和Accuracy数值一致(因为Accuracy本质就是全局正确样本数/总样本数,和micro平均的指标计算逻辑重合)
你的示例数据极度不平衡(5个样本属于类别2,仅1个样本属于类别1),micro平均会被多数类的贡献主导,进一步放大了这种指标重合的现象。
解决方法
创建指标时明确指定average参数,根据需求选择合适的平均方式:
average='macro':对每个类单独计算指标后取算术平均,不考虑类别样本量差异average='weighted':对每个类的指标按该类样本数量加权平均,适配不平衡数据集average=None:返回每个类的单独指标,不做平均处理
修改后的代码示例
import torch import torchmetrics from torchmetrics import MetricTracker, MetricCollection from torchmetrics import Accuracy, F1Score, Precision, Recall, CohenKappa num_classes = 3 # 指定average为weighted,适配不平衡数据场景 list_of_metrics = [ Accuracy(task="multiclass", num_classes=num_classes), F1Score(task="multiclass", num_classes=num_classes, average='weighted'), Precision(task="multiclass", num_classes=num_classes, average='weighted'), Recall(task="multiclass", num_classes=num_classes, average='weighted'), CohenKappa(task="multiclass", num_classes=num_classes) ] maximize_list=[True,True,True,True,True] metric_coll = MetricCollection(list_of_metrics) tracker = MetricTracker(metric_coll, maximize=maximize_list) pred = torch.Tensor([[0,.1,.5], # 预测为2 [0,.1,.5], # 预测为2 [0,.1,.5], # 预测为2 [0,.1,.5], # 预测为2 [0,.1,.5], # 预测为2 [0.9,.1,.5]]) # 预测为0 label = torch.Tensor([2,2,2,2,2,1]) tracker.increment() tracker.update(pred, label) for key, val in tracker.compute_all().items(): print(key,val)
修改后的输出
MulticlassAccuracy tensor([0.8333]) MulticlassF1Score tensor([0.7222]) MulticlassPrecision tensor([0.7222]) MulticlassRecall tensor([0.8333]) MulticlassCohenKappa tensor([0.4545])
可以看到,此时Precision、Recall、F1的结果不再与Accuracy一致,符合不平衡分类任务的预期指标表现。
如果需要分析单个类别的指标表现,可以将average设为None,此时会返回每个类的单独指标张量,例如:
F1Score(task="multiclass", num_classes=num_classes, average=None)
内容的提问来源于stack exchange,提问作者some_name.py
相关产品推荐
相关产品推荐

