如何将Detectron2训练的Faster-RCNN转为原生PyTorch模型?
问题:将Detectron2训练的Faster-RCNN转为原生PyTorch模型以适配GradCam
背景与模型加载情况
我用Detectron2训练了一个Faster-RCNN模型,权重保存为model.pth,同时配有对应的config.yml配置文件,可通过两种方式加载模型:
from detectron2.modeling import build_model from detectron2.checkpoint import DetectionCheckpointer from detectron2.config import get_cfg cfg = get_cfg() config_name = "config.yml" cfg.merge_from_file(config_name) # 方式1:通过DefaultPredictor加载 cfg.MODEL.WEIGHTS = './model.pth' model = DefaultPredictor(cfg) # 方式2:通过build_model + DetectionCheckpointer加载 model_ = build_model(cfg) model = DetectionCheckpointer(model_).load("./model.pth")
按照官方文档获取预测结果的方式:
import numpy as np from PIL import Image import torch image = np.array(Image.open('page4.jpg'))[:,:,::-1] # RGB转BGR格式 tensor_image = torch.from_numpy(image.copy()).permute(2, 0, 1) # 转为 (channels, H, W) 格式 with torch.no_grad(): output = model([{"image":tensor_image}])
查看模型类型的输出:
print(type(model)) print(type(model.model)) print(type(model.model.backbone))
输出结果:
<class 'detectron2.engine.defaults.DefaultPredictor'> <class 'detectron2.modeling.meta_arch.rcnn.GeneralizedRCNN'> <class 'detectron2.modeling.backbone.fpn.FPN'>
核心问题
需要用GradCam做模型可解释性分析,但GradCam要求使用原生PyTorch模型,此前尝试的两种方法均失败:
- 导入torchvision的Faster-RCNN加载权重:由于Detectron2和torchvision的Faster-RCNN层名称、结构尺寸不匹配,加载时报错
torch.save(model.model.state_dict(), "torch_weights.pth") from torchvision.models.detection import fasterrcnn_resnet50_fpn dummy = fasterrcnn_resnet50_fpn(pretrained=False, num_classes=1) dummy.load_state_dict(torch.load('./torch_weights.pth', map_location = 'cpu'))
- 自定义封装类:封装后的类无法直接访问
.backbone、.layers等底层属性,无法满足GradCam对中间层的访问需求
class TorchModel(torch.nn.Module): def __init__(self, model) -> None: super().__init__() self.model = model.model def forward(self, image): return self.model([{"image":image}])[0]['instances']
解决方案
方法1:直接基于Detectron2模型主体适配GradCam
Detectron2的GeneralizedRCNN本身就是继承自torch.nn.Module的原生PyTorch模型,只需做轻量封装即可适配GradCam的需求,同时保留对底层模块的访问能力:
from detectron2.config import get_cfg from detectron2.engine import DefaultPredictor import torch # 加载Detectron2模型 cfg = get_cfg() cfg.merge_from_file("config.yml") cfg.MODEL.WEIGHTS = './model.pth' predictor = DefaultPredictor(cfg) # 获取原生PyTorch模型主体(GeneralizedRCNN实例) detectron_model = predictor.model # 封装为GradCam兼容模型 class GradCamCompatibleModel(torch.nn.Module): def __init__(self, detectron_model): super().__init__() self.model = detectron_model # 指定GradCam要提取特征的目标层(以FPN的res5最后一层为例) self.target_layer = self.model.backbone.bottom_up.res5[-1] self.gradients = None self.activations = None # 注册钩子捕获梯度与激活值 def forward_hook(module, input, output): self.activations = output def backward_hook(module, grad_in, grad_out): self.gradients = grad_out[0] self.target_layer.register_forward_hook(forward_hook) self.target_layer.register_backward_hook(backward_hook) def forward(self, x): # 将输入转为Detectron2要求的格式 inputs = [{"image": x}] outputs = self.model(inputs) # 提取预测的最高类别分数,用于反向传播计算GradCam cls_scores = outputs[0]['instances'].scores return cls_scores.max()
封装后的模型既符合原生PyTorch模型规范,又能直接访问底层模块,可直接用于GradCam计算。
方法2:手动映射权重到torchvision模型(仅需使用torchvision模型时)
如果必须转为torchvision的fasterrcnn_resnet50_fpn,需手动映射Detectron2与torchvision的权重名称,示例如下:
import torch from torchvision.models.detection import fasterrcnn_resnet50_fpn from detectron2.config import get_cfg # 加载Detectron2权重 detectron_state_dict = torch.load('./model.pth', map_location='cpu')["model"] # 权重名称映射函数 def map_detectron_to_torchvision(detectron_dict): torchvision_dict = {} for k, v in detectron_dict.items(): new_k = k # 替换backbone层名称 new_k = new_k.replace("backbone.bottom_up", "backbone.body") new_k = new_k.replace("res2", "layer1") new_k = new_k.replace("res3", "layer2") new_k = new_k.replace("res4", "layer3") new_k = new_k.replace("res5", "layer4") # 保留roi_heads部分名称(结构一致时) torchvision_dict[new_k] = v return torchvision_dict # 初始化torchvision模型 cfg = get_cfg() cfg.merge_from_file("config.yml") torch_model = fasterrcnn_resnet50_fpn(pretrained=False, num_classes=cfg.MODEL.ROI_HEADS.NUM_CLASSES) # 加载映射后的权重,strict=False忽略不匹配的层 mapped_dict = map_detectron_to_torchvision(detectron_state_dict) torch_model.load_state_dict(mapped_dict, strict=False)
这种方法需根据实际模型结构调整映射规则,仅在必须使用torchvision模型时推荐使用。
内容的提问来源于stack exchange,提问作者Deshwal
相关产品推荐
相关产品推荐

