How to Get Both Losses and Predictions in Torchvision Faster R-CNN
If you are training an object detection model with PyTorch's torchvision.models.detection suite (such as fasterrcnn_resnet50_fpn_v2), you have likely run into an intentional design constraint: the model returns different outputs depending on its mode.
- In
model.train()mode, it requires ground-truth targets and returns a dictionary of losses (loss_classifier,loss_box_reg,loss_objectness,loss_rpn_box_reg). - In
model.eval()mode, it ignores targets and returns a list of prediction dictionaries containingboxes,labels, andscores.
Torchvision does this to avoid running expensive post-processing operations like Non-Maximum Suppression (NMS) during backward passes. However, when monitoring validation loss alongside mAP (Mean Average Precision), or when debugging false positives during training, getting both outputs from the same forward pass becomes crucial.
Below are two approaches: a complete single forward-pass solution via custom subclassing, and the standard two-step Torchvision validation workflow.
Method 1: Subclassing Faster R-CNN for a Single Forward Pass (Recommended)
The cleanest way to extract both losses and detections in a single pass is to extend the FasterRCNN architecture and slightly modify the internal RoIHeads behavior. When targets are supplied, we want the model to calculate losses and execute detection post-processing.
import torch
from torchvision.models.detection import fasterrcnn_resnet50_fpn_v2, FasterRCNN_ResNet50_FPN_V2_Weights
from torchvision.models.detection.roi_heads import RoIHeads
class DualOutputRoIHeads(RoIHeads):
"""Custom RoIHeads that compute both losses and detections simultaneously."""
def forward(self, features, proposals, image_shapes, targets=None):
if targets is not None:
# 1. Match targets to proposals to generate sampled proposals and compute losses
proposals, matched_idxs, labels, regression_targets = self.select_training_samples(proposals, targets)
box_features = self.box_roi_pool(features, proposals, image_shapes)
box_features = self.box_head(box_features)
class_logits, box_regression = self.box_predictor(box_features)
loss_classifier, loss_box_reg = self.box_loss(
class_logits, box_regression, labels, regression_targets
)
detector_losses = {"loss_classifier": loss_classifier, "loss_box_reg": loss_box_reg}
# 2. Run post-processing to obtain prediction boxes, labels, and scores
boxes, scores, labels = self.postprocess_detections(class_logits, box_regression, proposals, image_shapes)
num_images = len(boxes)
detections = []
for i in range(num_images):
detections.append({
"boxes": boxes[i],
"labels": labels[i],
"scores": scores[i],
})
return detections, detector_losses
else:
# Standard inference path
return super().forward(features, proposals, image_shapes, targets)
def get_dual_output_fasterrcnn(num_classes=91, weights=None):
# Initialize the base model
model = fasterrcnn_resnet50_fpn_v2(weights=weights)
# Replace standard roi_heads with our DualOutputRoIHeads
custom_roi_heads = DualOutputRoIHeads(
box_roi_pool=model.roi_heads.box_roi_pool,
box_head=model.roi_heads.box_head,
box_predictor=model.roi_heads.box_predictor,
fg_iou_thresh=model.roi_heads.proposal_matcher.high_threshold,
bg_iou_thresh=model.roi_heads.proposal_matcher.low_threshold,
batch_size_per_image=model.roi_heads.fg_bg_sampler.batch_size_per_image,
positive_fraction=model.roi_heads.fg_bg_sampler.positive_fraction,
bbox_reg_weights=model.roi_heads.box_coder.weights,
score_thresh=model.roi_heads.score_thresh,
nms_thresh=model.roi_heads.nms_thresh,
detections_per_img=model.roi_heads.detections_per_img,
)
model.roi_heads = custom_roi_heads
# Modify the top-level forward pass
original_forward = model.forward
def forward_both(images, targets=None):
if model.training and targets is None:
raise ValueError("In training mode, targets should be passed")
# Standard GeneralizedRCNN transform and feature extraction
original_image_sizes = [img.shape[-2:] for img in images]
images, targets = model.transform(images, targets)
features = model.backbone(images.tensors)
# Region Proposal Network (RPN)
proposals, proposal_losses = model.rpn(images, features, targets)
# RoI Heads (Computes both losses and detections)
detections, detector_losses = model.roi_heads(features, proposals, images.image_sizes, targets)
detections = model.transform.postprocess(detections, images.image_sizes, original_image_sizes)
losses = {}
losses.update(detector_losses)
losses.update(proposal_losses)
return losses, detections
model.forward = forward_both
return model
How to use it:
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
model = get_dual_output_fasterrcnn(weights=FasterRCNN_ResNet50_FPN_V2_Weights.DEFAULT).to(device)
model.train()
images = [torch.rand(3, 400, 400).to(device)]
targets = [{
"boxes": torch.tensor([[50.0, 50.0, 150.0, 150.0]]).to(device),
"labels": torch.tensor([1], dtype=torch.int64).to(device)
}]
losses, detections = model(images, targets)
print("Losses:", losses.keys())
print("Predictions:", detections[0].keys())
Method 2: The Non-Invasive Validation Trick (Without Modifying the Model)
If you prefer not to touch the internal mechanics of Torchvision, note that you can calculate validation losses simply by keeping the model in model.train() mode inside a torch.no_grad() block:
def validate_epoch(model, val_loader, device):
total_val_loss = 0.0
with torch.no_grad():
for images, targets in val_loader:
images = [img.to(device) for img in images]
targets = [{k: v.to(device) for k, v in t.items()} for t in targets]
# 1. Get Losses by calling model in train() mode
model.train()
loss_dict = model(images, targets)
batch_loss = sum(loss for loss in loss_dict.values())
total_val_loss += batch_loss.item()
# 2. Get Predictions by toggling to eval() mode
model.eval()
predictions = model(images)
# Update your mAP metric (e.g. torchmetrics.detection.MeanAveragePrecision)
# metric.update(predictions, targets)
return total_val_loss / len(val_loader)
Performance Note: Avoid NMS on Every Training Step
While extracting detections during training is useful for debugging and tracking visual qualitative results, running post-processing (NMS and bounding-box sorting) on every batch significantly slows down training speed.
For production workflows, it is best to calculate predictions only periodically (e.g., at the end of each epoch or every N batches) using Method 2, or apply Method 1 conditionally based on a parameter flag.