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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 06:47:18