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

Python中如何重写FasterRCNN父类实例并自定义子类?

正确实现基于FasterRCNN工厂方法的自定义子类

你原来的实现存在核心问题:通过__new__返回fasterrcnn_resnet50_fpn()生成的实例时,这个实例本质是FasterRCNN类的对象,而非你的MyDetector子类实例。这会导致__init__代码不会执行,自定义的some_func也无法调用,因为对象根本不属于MyDetector类。

下面提供两种可行的实现方式:

方式一:组合工厂模型(推荐,简单高效)

直接复用工厂方法生成的模型,将其作为自定义类的内部属性,通过方法代理让自定义类拥有原模型的所有功能,同时添加自定义方法:

from torchvision.models.detection import fasterrcnn_resnet50_fpn, FasterRCNN_ResNet50_FPN_Weights
from torchvision.models.detection.faster_rcnn import FastRCNNPredictor

class MyDetector:
    def __init__(self):
        # 用工厂方法创建预训练的FasterRCNN模型
        self.model = fasterrcnn_resnet50_fpn(weights=FasterRCNN_ResNet50_FPN_Weights.DEFAULT)
        
        # 修改预测器适配自定义类别数
        num_features_in = self.model.roi_heads.box_predictor.cls_score.in_features
        self.model.roi_heads.box_predictor = FastRCNNPredictor(num_features_in, num_classes=2)
    
    def some_func(self):
        # 自定义业务方法,示例:打印模型信息
        print(f"当前模型类别数:{self.model.roi_heads.box_predictor.cls_score.out_features}")
    
    # 代理原模型的所有方法/属性,让MyDetector实例可以直接调用原模型的方法
    def __getattr__(self, name):
        return getattr(self.model, name)

使用示例:

detector = MyDetector()
detector.some_func()  # 调用自定义方法
outputs = detector(input_images)  # 直接调用原模型的推理方法

方式二:真正继承FasterRCNN类

如果需要严格的继承关系,可以复用工厂方法的组件参数,手动初始化父类:

from torchvision.models.detection import FasterRCNN, FasterRCNN_ResNet50_FPN_Weights
from torchvision.models.detection.faster_rcnn import FastRCNNPredictor
from torchvision.models.detection.backbone_utils import resnet_fpn_backbone

class MyDetector(FasterRCNN):
    def __init__(self, num_classes=2):
        # 复用工厂方法的backbone(预训练的ResNet50+FPN)
        backbone = resnet_fpn_backbone(
            'resnet50', 
            weights=FasterRCNN_ResNet50_FPN_Weights.DEFAULT.backbone
        )
        
        # 调用父类FasterRCNN的构造函数,传入工厂方法的默认参数
        super().__init__(
            backbone=backbone,
            num_classes=91,  # 预训练模型默认类别数,后续替换预测器
            min_size=800,
            max_size=1333,
            rpn_nms_thresh=0.7,
            box_score_thresh=0.05,
            box_nms_thresh=0.5,
            # 其他参数可参考fasterrcnn_resnet50_fpn的源码默认值
        )
        
        # 替换预测器为自定义类别数
        num_features_in = self.roi_heads.box_predictor.cls_score.in_features
        self.roi_heads.box_predictor = FastRCNNPredictor(num_features_in, num_classes)
    
    def some_func(self):
        print("执行自定义操作")

这种方式下,MyDetector是真正的FasterRCNN子类,所有父类方法都可以直接使用,也能重写父类方法(如forward)进行深度定制。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.13 19:01:01