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不一致,会触发变量未定义错误。
修复方案
二值图像分割场景下,按以下逻辑调整代码即可:
- 初始化
Precision时显式指定task="binary",可按需设置二值化阈值(默认0.5),二分类场景不需要手动传入num_classes参数。 - 计算指标前先对模型输出做sigmoid激活,将logits转为0~1区间的概率值,再挤压掉长度为1的通道维度,把预测、标签的形状从
(batch_size, 1, h, w)转为(batch_size, h, w)。 - 修正原代码的变量名笔误。
修正后的核心代码如下:
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
相关产品推荐
相关产品推荐

