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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.03 03:45:46