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
相关产品推荐
相关产品推荐

