自定义PyTorch FasterRCNN前向传播异常:detections为空如何修复?
自定义FasterRCNN前向传播导致detections为空的解决方法
问题核心原因
你手动实现的前向传播逻辑忽略了原FasterRCNN模型的关键预处理和模式区分逻辑,主要问题包括:
- 跳过了模型内置的图像缩放、padding预处理步骤,导致输入特征与RPN、RoI Heads的预期不匹配
- Backbone输入未使用经过正确预处理的张量
- 未区分训练/评估模式,评估阶段缺少检测结果的后处理步骤
修正后的代码
from torchvision.models.detection.image_list import ImageList from collections import OrderedDict import torch import torch.nn as nn from torchvision.models.detection import fasterrcnn_resnet50_fpn from torchvision.models.detection.faster_rcnn import FastRCNNPredictor from torchvision.models.detection.faster_rcnn import FasterRCNN_ResNet50_FPN_Weights class CustomFasterRCNNResNet50FPN(nn.Module): def __init__(self, num_classes, **kwargs): super().__init__() # 加载预训练模型 self.model = fasterrcnn_resnet50_fpn(weights=FasterRCNN_ResNet50_FPN_Weights.DEFAULT) # 替换分类器 in_features = self.model.roi_heads.box_predictor.cls_score.in_features self.model.roi_heads.box_predictor = FastRCNNPredictor(in_features, num_classes) def forward(self, images, targets=None): # 1. 复用原模型的图像预处理逻辑,自动处理缩放、padding image_list = self.model.transform(images, targets) # 2. Backbone输入使用预处理后的张量 features = self.model.backbone(image_list.tensors) if isinstance(features, torch.Tensor): features = OrderedDict([("0", features)]) # 3. 调用RPN和RoI Heads proposals, proposal_losses = self.model.rpn(image_list, features, targets) detections, detector_losses = self.model.roi_heads(features, proposals, image_list.image_sizes, targets) # 4. 区分训练/评估模式的返回值 if self.training: losses = {} losses.update(detector_losses) losses.update(proposal_losses) return losses else: # 评估阶段:将检测框映射回原始图像尺寸 return self.model.transform.postprocess(detections, image_list.image_sizes, [img.shape[-2:] for img in images])
关键修改说明
- 复用模型内置transform:原模型的
transform模块会自动将输入图像调整到符合模型要求的尺寸,并添加padding,这是手动stack图像无法替代的核心步骤 - 正确传递Backbone输入:使用
image_list.tensors作为Backbone的输入,确保特征提取的正确性 - 模式区分与后处理:评估阶段调用
postprocess将检测框坐标映射回原始图像尺寸,避免因坐标不匹配导致的空检测结果 - 保留原模型逻辑对齐:严格遵循原GeneralizedRCNN的前向传播流程,确保各模块输入输出匹配
额外检查点
- 确认
num_classes设置正确:若包含背景类,需设置为实际类别数+1(比如5类目标需设为6) - 评估时切换模型模式:调用
model.eval(),避免训练模式下的RPN行为干扰检测结果输出
内容的提问来源于stack exchange,提问作者morepenguins
相关产品推荐
相关产品推荐

