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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.28 09:06:03