使用Torchvision RetinaNet时输入数据结构引发分类头损失报错如何解决
Torchvision RetinaNet分类损失计算IndexError问题解决
报错根因
- 核心问题出在类别编号和
num_classes参数的匹配规则不符合Torchvision检测模型的约定:
Torchvision的所有目标检测模型(包括RetinaNet)的num_classes参数需要包含背景类,你当前要检测1类正目标,设置num_classes=1的情况下,分类头的输出通道仅支持索引0的类别,但是你构造的标签全部取值为1,访问分类输出的第1位时就触发了维度越界。 - 额外小问题:你的代码中使用了
np.ones但未导入numpy包,也会触发导入错误。 - 至于仅分类头报错、回归头正常的原因:回归损失计算只用到框坐标,不需要读取类别标签,自然不会触发索引越界。
修复方案
方案1(符合官方约定,更推荐)
将模型初始化的num_classes参数改为2,对应「1个正类别+1个背景类」,不需要修改标签取值:
model = RetinaNet(backbone, num_classes=2, # 原先是1,修改为2 anchor_generator=anchor_generator)
方案2(保持num_classes=1)
将生成标签的代码改为取值全为0即可:
labels = torch.LongTensor(np.zeros(len(bboxes), dtype=int))
补充修复
在代码开头添加numpy导入:
import numpy as np
修复后验证
修改后运行代码即可正常输出损失字典,不会再触发IndexError。
内容的提问来源于stack exchange,提问作者asi212
相关产品推荐
相关产品推荐

