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

torchmetrics计算Precision报错:隐含类别数与num_classes不匹配

报错原因
  • 新版torchmetrics的Precision指标必须显式指定task参数,未指定时默认按多分类(multiclass)逻辑做形状校验,和二值分割的输入格式要求不匹配。
  • 多分类模式下要求标签形状为(batch_size, h, w)(无单独类别通道维度),标签值为0到num_classes-1的整数索引。你传入的预测、标签都保留了channels=1的维度,形状为(batch_size, 1, h, w),和你设置的num_classes=1的校验规则冲突,直接触发报错。
  • 额外问题:你传入的outputs是模型输出的原始logits,未经过sigmoid激活和二值化处理,就算不报错计算出的精度结果也完全错误;另外原代码存在笔误,outputs = model(input)中传入的变量名和前面定义的inputs不一致,会触发变量未定义错误。
修复方案

二值图像分割场景下,按以下逻辑调整代码即可:

  1. 初始化Precision时显式指定task="binary",可按需设置二值化阈值(默认0.5),二分类场景不需要手动传入num_classes参数。
  2. 计算指标前先对模型输出做sigmoid激活,将logits转为0~1区间的概率值,再挤压掉长度为1的通道维度,把预测、标签的形状从(batch_size, 1, h, w)转为(batch_size, h, w)。
  3. 修正原代码的变量名笔误。

修正后的核心代码如下:

inputs, labels = batch
outputs = model(inputs) # 修复原代码input变量名笔误
loss = self.loss_function(outputs, labels)

# 初始化二值分割精度指标
prec_cal = torchmetrics.Precision(task="binary", threshold=0.5)
# 处理预测值:sigmoid转概率 + 去除通道维
preds = torch.sigmoid(outputs).squeeze(dim=1)
# 处理标签:去除通道维 + 转整数类型
targets = labels.squeeze(dim=1).long()
prec = prec_cal(preds, targets)

如果你的torchmetrics版本较低不支持task参数,可额外将num_classes=2传入,同时把预测结果通过阈值转为0/1的类别索引,再去掉通道维传入,不过更推荐升级到新版使用task="binary"的标准写法。

内容的提问来源于stack exchange,提问作者jumpincrane

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 02:27:20