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

torchmetrics与PyTorch Lightning搭配使用时多分类指标结果一致问题

问题原因

你遇到的指标输出完全一致是参数默认设置导致的必然结果:
你使用的F1、Precision、Recall三个指标默认采用了average="micro"的计算策略。micro平均的逻辑是先统计所有类别的总真阳性(TP)、总假阳性(FP)、总假阴性(FN),再基于全局统计值计算指标。
在多分类场景下,每个样本仅对应一个真实类别,此时:

  • 总TP + 总FP = 全部样本数量(所有预测结果的总数)
  • 总TP + 总FN = 全部样本数量(所有真实标签的总数)
    因此 micro精度 = 总TP/(总TP+总FP) = 总TP/总样本数,micro召回 = 总TP/(总TP+总FN) = 总TP/总样本数,micro F1是两者的调和平均自然也等于该值,而多分类准确率本身就是正确预测数/总样本数,所以四个指标输出完全一致。
解决方法

根据你的实际需求修改指标初始化时的average参数即可,新版torchmetrics还需要明确指定任务类型和类别数避免警告:

  • 如需获取每个类别的单独指标:设置average=None,输出会是和类别数等长的张量,每个位置对应对应类别的指标值
  • 如需不考虑类别样本量的宏平均:设置average="macro",会先计算每个类别的指标再取算术平均
  • 如需按类别样本量加权的加权平均:设置average="weighted",会先计算每个类别的指标再按对应类别的真实样本占比加权求和

修改后的参考代码如下:

import torch
import torchmetrics

# 示例使用macro平均,可根据需求替换average参数
metric_acc = torchmetrics.Accuracy(task="multiclass", num_classes=5)
metric_f1 = torchmetrics.F1(task="multiclass", num_classes=5, average="macro")
metric_pre = torchmetrics.Precision(task="multiclass", num_classes=5, average="macro")
metric_rec = torchmetrics.Recall(task="multiclass", num_classes=5, average="macro")

n_batches = 3
for i in range(n_batches):
    # simulate a classification problem
    preds = torch.randn(10, 5).softmax(dim=-1)
    target = torch.randint(5, (10,))

    acc = metric_acc(preds, target)
    f1 = metric_f1(preds, target)
    pre = metric_pre(preds, target)
    rec = metric_rec(preds, target)
    print(f"Accuracy on batch {i}: {acc}")
    print(f"F1 score on batch {i}: {f1}")
    print(f"pre score on batch {i}: {pre}")
    print(f"rec score on batch {i}: {rec}")
    print('-' * 20)


acc = metric_acc.compute()
f1 = metric_f1.compute()
pre = metric_pre.compute()
rec = metric_rec.compute()
print(f"Accuracy on all data: {acc}")
print(f"f1 score on all data: {f1}")
print(f"pre score on all data: {pre}")
print(f"rec score on all data: {rec}")

内容的提问来源于stack exchange,提问作者AI-bobobo

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 06:54:02