使用torchmetrics计算多标签分类精确率与召回率的参数问题
多标签分类任务Precision与Recall计算方案
TorchMetrics提供了可直接用于多标签分类(单样本可归属多个类别)场景的精确率、召回率计算接口,你之前得到不符合预期的返回值,是参数配置错误导致的。
错误原因
你代码中使用的mdmc_average参数是为多维多分类(multi-dimensional multi-class)场景设计的,不适用于普通多标签分类场景。错误的参数组合让模块误判了任务类型,执行了不符合多标签逻辑的计算,才会出现真阳性为0时精确率返回0.7的异常结果。
正确使用方式
多标签场景需要显式指定task="multilabel"声明任务类型,无需配置mdmc_average参数,参考代码如下:
import torch from torchmetrics import Precision, Recall target = torch.tensor([ [0, 0, 1, 1, 0], # 样本1属于类别2、3(零索引) [0, 0, 1, 0, 0], # 样本2属于类别2(零索引) ]) preds = torch.tensor([ [0, 0, 0, 0, 0], # 样本1预测无所属类别 [0, 0, 0, 0, 0], # 样本2预测无所属类别 ]) # 初始化多标签任务指标 precision_metric = Precision( task="multilabel", num_labels=5, average="samplewise" ) recall_metric = Recall( task="multilabel", num_labels=5, average="samplewise" ) print(precision_metric(preds, target)) # 输出tensor(0.),符合预期 print(recall_metric(preds, target)) # 输出tensor(0.),符合预期
关键参数说明
task="multilabel":多标签分类场景必传参数,显式声明任务类型,避免模块误判计算逻辑num_labels:多标签任务的总类别数,替代多分类场景使用的num_classes参数average:控制指标聚合方式,传入samplewise时会先计算每个样本单独的指标值,再对所有样本求平均,匹配逐样本计算的需求
边界场景说明:如果遇到样本预测全为负例、真实标签也全为负例的情况,默认会将该样本的精确率、召回率记为0,可通过
zero_division参数自定义该场景下的返回值。
内容的提问来源于stack exchange,提问作者RR_28023
相关产品推荐
相关产品推荐

