How RF-DETR Infers the Model Class from a Checkpoint

RF-DETR reconstructs the exact model class automatically by inspecting the hyper_parameters dictionary stored in the Lightning checkpoint, specifically the model_cfg and model_name keys, then resolving the class via an internal registry.

When loading a saved checkpoint, the RF-DETR framework must determine which specific model architecture—such as RFDETRSmall or RFDETRSegMedium—generated the file. Rather than requiring manual specification, the library embeds sufficient metadata during training to enable automatic class inference. This process centers on the on_load_checkpoint hook in src/rfdetr/training/module_model.py, which orchestrates legacy format detection, configuration extraction, and class resolution.

The Checkpoint Loading Pipeline

The entry point for this logic is RFDETRModelModule.on_load_checkpoint, invoked automatically by PyTorch Lightning when load_from_checkpoint is called. The method implements a five-step pipeline to ensure the correct class is instantiated regardless of whether the checkpoint originated from the modern training loop or a legacy .pth file.

1. Legacy Format Detection

First, the method checks for the legacy_checkpoint_format boolean flag. If present and True, the checkpoint was produced by the older engine-based training loop and converted via convert_legacy_checkpoint in src/rfdetr/training/checkpoint.py. The routine knows to strip the "model." prefix from state dict keys before processing.

2. Configuration Extraction

The checkpoint's hyper_parameters dictionary contains the keys "model_cfg" (a dictionary describing architecture parameters) and optionally "model_name" (a string identifier). The loader extracts these values to determine what class to construct.

model_cfg = checkpoint["hyper_parameters"]["model_cfg"]
model_name = checkpoint["hyper_parameters"].get("model_name")

3. Registry Resolution

RF-DETR maintains an internal MODEL_REGISTRY mapping model names to concrete Python classes defined under src/rfdetr/models/. The loader resolves the class by looking up model_name in this registry, falling back to model_cfg["model_type"] if the name is absent.

from rfdetr.platform.models import MODEL_REGISTRY
ModelCls = MODEL_REGISTRY[model_name or model_cfg["model_type"]]

4. Model Instantiation

With the resolved class, the method creates a fresh instance using the saved hyper-parameters (number of classes, backbone settings, etc.). It then loads the processed state dict onto this instance.

5. EMA Handling

If the checkpoint contains an ema_state_dict or legacy_ema_state_dict, the routine stores it for optional restoration but does not allow it to interfere with class inference.

Practical Code Examples

The following examples demonstrate how RF-DETR's automatic class inference works in practice.

Loading a Checkpoint Directly

from rfdetr.training.module_model import RFDETRModelModule

# Automatic class inference happens inside load_from_checkpoint

module = RFDETRModelModule.load_from_checkpoint("path/to/checkpoint.ckpt")

# module.model is now the correct class instance

print(type(module.model))

# <class 'rfdetr.models.detr_small.RFDETRSmall'>

Inspecting Checkpoint Metadata Manually

import torch

ckpt = torch.load("path/to/checkpoint.ckpt", map_location="cpu")
print(ckpt["hyper_parameters"]["model_name"])

# Output: "rfdetr-small"

Converting and Loading Legacy Checkpoints

from rfdetr.training.checkpoint import convert_legacy_checkpoint

# Convert old .pth to new format

convert_legacy_checkpoint("legacy.ckpt.pth", "new.ckpt")

# Load with automatic class detection

module = RFDETRModelModule.load_from_checkpoint("new.ckpt")

Source File Reference Table

File Purpose
src/rfdetr/training/module_model.py Contains RFDETRModelModule.on_load_checkpoint, the core logic for class inference and checkpoint loading.
src/rfdetr/training/checkpoint.py Provides convert_legacy_checkpoint to normalize old checkpoints for the new loader.
tests/training/test_module_model.py Validates automatic class inference, legacy format handling, and EMA restoration.

Summary

  • Metadata-driven inference: RF-DETR stores model_cfg and model_name in the checkpoint's hyper_parameters to eliminate manual class specification.
  • Legacy compatibility: The legacy_checkpoint_format flag ensures older checkpoints are handled correctly by stripping the "model." prefix.
  • Registry pattern: Class resolution uses an internal MODEL_REGISTRY mapping string identifiers to concrete model classes.
  • Seamless loading: Users call RFDETRModelModule.load_from_checkpoint() without knowing the specific model class in advance.

Frequently Asked Questions

How does RF-DETR handle checkpoints from older versions of the codebase?

RF-DETR detects legacy checkpoints via the legacy_checkpoint_format flag set during conversion. The on_load_checkpoint method in src/rfdetr/training/module_model.py checks this flag and strips the "model." state dict prefix before class inference proceeds.

What happens if the model_name key is missing from the checkpoint?

If model_name is absent, the loader falls back to model_cfg["model_type"] to determine the correct class. Both keys map to entries in the MODEL_REGISTRY, ensuring robust class resolution even with partial metadata.

Can I load a checkpoint without using the RFDETRModelModule wrapper?

While possible, manual loading requires replicating the logic in on_load_checkpoint: extracting hyper_parameters, resolving the class via MODEL_REGISTRY, and handling legacy prefixes. Using the module's load_from_checkpoint method is strongly recommended to ensure consistent behavior.

Where is the model architecture actually defined when loading?

The concrete classes (e.g., RFDETRSmall) reside in src/rfdetr/models/. The checkpoint only stores the string identifier and configuration; the MODEL_REGISTRY in the platform layer maps these to the actual Python classes during instantiation.

Have a question about this repo?

These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →