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

Detectron2自定义FastRCNNOutputLayers报错:参数数量不匹配求助

解决Detectron2中CustomROIHeads初始化参数不匹配问题

错误原因分析

报错TypeError: CustomROIHeads.__init__() takes 2 positional arguments but 3 were given本质是两个问题:

  • Detectron2的build_roi_heads函数会向ROIHeads类传入cfg和input_shape两个位置参数,加上self共3个参数,但你定义的CustomROIHeads.__init__只接受2个参数,数量不匹配
  • 直接在__init__中实例化CustomRCNNOutput时缺少input_shape参数,不符合其构造函数要求

修正方案

步骤1:给自定义输出层注册到Detectron2注册器

Detectron2通过注册器管理组件,必须将CustomRCNNOutput注册到FAST_RCNN_OUTPUT_LAYERS_REGISTRY:

from detectron2.modeling.roi_heads.fast_rcnn import FastRCNNOutputLayers, FAST_RCNN_OUTPUT_LAYERS_REGISTRY
from detectron2.structures import Instances
from detectron2.utils.events import _log_classification_stats
from torch import cat

@FAST_RCNN_OUTPUT_LAYERS_REGISTRY.register()
class CustomRCNNOutput(FastRCNNOutputLayers):
    def __init__(self, cfg, input_shape):
        super().__init__(cfg, input_shape)
    
    def losses(self, predictions, proposals):
        scores, proposal_deltas = predictions

        gt_classes = (
            cat([p.gt_classes for p in proposals], dim=0) if len(proposals) else predictions[0].new_empty(0)
        )
        _log_classification_stats(scores, gt_classes)

        if len(proposals):
            proposal_boxes = cat([p.proposal_boxes.tensor for p in proposals], dim=0)
            assert not proposal_boxes.requires_grad, "Proposals should not require gradients!"
            gt_boxes = cat(
                [(p.gt_boxes if p.has("gt_boxes") else p.proposal_boxes).tensor for p in proposals],
                dim=0,
            )
        else:
            proposal_boxes = gt_boxes = predictions[0].new_empty((0, 4))

        if self.use_sigmoid_ce:
            loss_cls = self.sigmoid_cross_entropy_loss(scores, gt_classes)
        else:
            # 替换为你的自定义损失函数
            loss_cls = MY_CUSTOM_LOSS(scores, gt_classes, self.num_classes)

        losses = {
            "loss_cls": loss_cls,
            "loss_box_reg": self.box_reg_loss(
                proposal_boxes, gt_boxes, proposal_deltas, gt_classes
            ),
        }
        return {k: v * self.loss_weight.get(k, 1.0) for k, v in losses.items()}

步骤2:修正CustomROIHeads的构造函数与box_predictor配置

修改CustomROIHeads的__init__签名匹配父类,并通过配置指定使用自定义输出层:

from detectron2.modeling.roi_heads import StandardROIHeads, ROI_HEADS_REGISTRY

@ROI_HEADS_REGISTRY.register()
class CustomROIHeads(StandardROIHeads):
    def __init__(self, cfg, input_shape, **kwargs):
        # 克隆配置并修改box_predictor为自定义类
        cfg = cfg.clone()
        cfg.MODEL.ROI_HEADS.BOX_PREDICTOR = "CustomRCNNOutput"
        # 调用父类构造函数,自动创建自定义box_predictor
        super().__init__(cfg, input_shape, **kwargs)

可选方案:重写from_config方法(不修改全局配置)

如果不想修改cfg,也可以通过重写from_config方法直接替换box_predictor:

@ROI_HEADS_REGISTRY.register()
class CustomROIHeads(StandardROIHeads):
    @classmethod
    def from_config(cls, cfg, input_shape):
        # 获取父类的默认配置
        ret = super().from_config(cfg, input_shape)
        # 替换box_predictor为自定义实例,传入box_head的输出形状
        ret["box_predictor"] = CustomRCNNOutput(cfg, ret["box_head"].output_shape)
        return ret

    def __init__(self, cfg, input_shape, **kwargs):
        # 保持构造函数签名与父类一致
        super().__init__(cfg, input_shape, **kwargs)

步骤3:在训练配置中指定使用CustomROIHeads

在你的训练脚本中添加:

cfg.MODEL.ROI_HEADS.NAME = "CustomROIHeads"

关键说明

  • 自定义组件的构造函数签名必须与Detectron2基类一致,否则会出现参数不匹配错误
  • 所有自定义组件都需要注册到对应Detectron2注册器,框架才能自动发现并实例化
  • 不要在__init__中直接实例化依赖其他组件输出形状的模块(如box_predictor),应通过注册器或from_config方法延迟创建

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 15:15:38