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

