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_cfgandmodel_namein the checkpoint'shyper_parametersto eliminate manual class specification. - Legacy compatibility: The
legacy_checkpoint_formatflag ensures older checkpoints are handled correctly by stripping the"model."prefix. - Registry pattern: Class resolution uses an internal
MODEL_REGISTRYmapping 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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →