How to Train RF-DETR with a Custom Dataset: A Complete Guide
To train RF-DETR on a custom dataset, use RFDETRMedium().train(dataset_dir="./my_dataset") for a one-line quick start, or instantiate RFDETRDataModule and RFDETRModelModule with PyTorch Lightning for full control over the training pipeline.
The roboflow/rf-detr repository provides a real-time transformer-based detection framework built on PyTorch Lightning. Whether you are fine-tuning for object detection, instance segmentation, or key-point estimation, the training workflow automatically handles dataset parsing, augmentation, and checkpoint management with minimal boilerplate.
Quick Start with the One-Line API
The fastest way to train RF-DETR with a custom dataset is through the high-level train() method available on every model class. This single call constructs the data module, model module, and Lightning trainer internally.
from rfdetr import RFDETRMedium
model = RFDETRMedium() # Downloads COCO-pretrained weights automatically
model.train(
dataset_dir="./my_dataset",
epochs=100,
batch_size="auto", # Probes GPU memory to set a practical batch size
lr=1e-4,
output_dir="./rf_detr_output",
)
The batch_size="auto" parameter probes available GPU memory and selects a practical batch size automatically. Checkpoints are saved in the specified output_dir, including checkpoint_best_ema.pth and checkpoint_best_regular.pth when using EMA weighting.
Dataset Format Requirements
RF-DETR accepts both COCO and YOLO formatted datasets without requiring explicit format declarations. The RFDETRDataModule in src/rfdetr/training/module_data.py automatically detects the structure and builds the appropriate PyTorch datasets.
COCO Format
Place train/_annotations.coco.json inside your dataset root. Include valid/_annotations.coco.json for validation splits.
YOLO Format
Include a data.yaml file alongside train/images, train/labels, valid/images, and valid/labels directories.
Key-Point Preview Datasets
For pose estimation, use COCO key-point JSON or Ultralytics YOLO pose datasets that define kpt_shape in data.yaml. The infer_coco_keypoint_schema utility automatically parses the key-point schema from these files.
The Core Training Architecture
According to the roboflow/rf-detr source code, the training pipeline consists of three interchangeable components managed by PyTorch Lightning:
RFDETRDataModule (src/rfdetr/training/module_data.py)
The data module parses dataset directories, infers class counts and key-point schemas, and applies augmentations. It supports three augmentation backends via the augmentation_backend parameter:
"cpu"– Albumentations transforms"gpu"– Kornia transforms"torchvision"– Default TorchVision transforms
RFDETRModelModule (src/rfdetr/training/module_model.py)
Located in src/rfdetr/training/module_model.py, this module wraps the RF-DETR backbone, classification head, and task-specific heads (bounding box, segmentation mask, or key-point). It implements the training_step, validation_step, and loss calculation logic.
PyTorch Lightning Trainer
The trainer coordinates epochs, validation loops, and callbacks. It includes EMA (Exponential Moving Average) weighting, early stopping, and automatic checkpointing out-of-the-box.
Training for Different Tasks
RF-DETR provides specialized model classes for each computer vision task. All share the same train() signature but load different heads and loss functions.
Object Detection
Use RFDETRMedium or RFDETRLarge for standard bounding box detection:
from rfdetr import RFDETRMedium
model = RFDETRMedium()
model.train(
dataset_dir="./my_coco_dataset",
epochs=80,
batch_size="auto",
lr=1e-4,
output_dir="./output_detection",
)
Instance Segmentation
Use RFDETRSegMedium or RFDETRSegLarge for pixel-level segmentation:
from rfdetr import RFDETRSegMedium
model = RFDETRSegMedium()
model.train(
dataset_dir="./my_yolo_dataset",
epochs=120,
batch_size="auto",
lr=5e-5,
output_dir="./output_segmentation",
)
Key-Point Preview
Use RFDETRKeypointPreview for custom pose datasets. You must infer the schema and pass key-point specific parameters:
from pathlib import Path
from rfdetr import RFDETRKeypointPreview
from rfdetr.datasets._keypoint_schema import infer_coco_keypoint_schema
DATASET = Path("./my_keypoint_coco")
schema = infer_coco_keypoint_schema(DATASET / "train" / "_annotations.coco.json")
model = RFDETRKeypointPreview(
num_classes=len(schema.class_names),
num_keypoints_per_class=schema.num_keypoints_per_class,
)
model.train(
dataset_dir=str(DATASET),
class_names=schema.class_names,
keypoint_oks_sigmas=schema.keypoint_oks_sigmas,
epochs=50,
batch_size=8,
lr=2e-5,
output_dir="./output_keypoints",
)
Advanced Custom Training
When you need custom callbacks, multi-GPU DDP strategies, or alternative loggers (Weights & Biases, MLflow), instantiate the Lightning components directly rather than using the one-line API.
from rfdetr.training.module_data import RFDETRDataModule
from rfdetr.training.module_model import RFDETRModelModule
from pytorch_lightning import Trainer
dm = RFDETRDataModule(
dataset_dir="./my_dataset",
batch_size=4,
augmentation_backend="gpu", # Forces Kornia augmentations on GPU
)
model = RFDETRModelModule(
num_classes=dm.num_classes,
num_keypoints=dm.num_keypoints
)
trainer = Trainer(
max_epochs=30,
accelerator="gpu",
devices=2, # Multi-GPU DDP
callbacks=[MyEarlyStoppingCallback(patience=5)],
logger=my_wandb_logger,
)
trainer.fit(model, dm)
This approach gives you full control over the training loop while retaining the dataset parsing logic and model architecture defined in module_data.py and module_model.py.
Monitoring and Checkpoints
Training metrics are visualized using helpers in src/rfdetr/visualize/training.py, which generates loss curves and AP charts. By default, the system logs to TensorBoard, but you can install additional loggers via pip install "rfdetr[train,loggers]".
EMA weights are used for the "best" checkpoint by default. Disable this behavior by passing use_ema=False to model.train(). To resume training from an interruption, pass the checkpoint path to the resume parameter:
model.train(
dataset_dir="./my_dataset",
resume="./rf_detr_output/checkpoint_last.pth",
epochs=100,
)
Summary
- Use
RFDETRMedium().train()or equivalent model classes for one-line training on COCO or YOLO datasets. - The training stack relies on
RFDETRDataModule(src/rfdetr/training/module_data.py) for data loading andRFDETRModelModule(src/rfdetr/training/module_model.py) for model logic. - Automatic format detection handles COCO JSON and YOLO directory structures without explicit configuration.
- EMA weighting, early stopping, and automatic checkpointing are enabled by default.
- For advanced use cases, manually instantiate PyTorch Lightning components to customize callbacks, loggers, and distributed training strategies.
Frequently Asked Questions
What dataset formats does RF-DETR support for training?
RF-DETR automatically detects and supports COCO JSON format (with _annotations.coco.json files) and YOLO format (with data.yaml and image/label directories). For key-point preview training, it also accepts COCO key-point JSON or YOLO pose datasets that define kpt_shape in their configuration files. The RFDETRDataModule handles format inference internally.
How do I resume training from a checkpoint?
Pass the checkpoint path to the resume parameter in the train() method: model.train(resume="./output/checkpoint_last.pth", ...). The Lightning trainer will restore model weights, optimizer states, and the current epoch count, allowing you to continue training seamlessly from where you left off.
Can I train RF-DETR on multiple GPUs?
Yes. While the one-line API uses single-GPU training by default, you can enable multi-GPU training by using the Custom Training API. Instantiate RFDETRDataModule and RFDETRModelModule manually, then pass devices=2 (or more) and strategy="ddp" to the PyTorch Lightning Trainer object.
What is the difference between the "best" and "last" checkpoints?
The "best" checkpoint (checkpoint_best_ema.pth) uses Exponential Moving Average weights that typically generalize better than the regular weights. The "last" checkpoint (checkpoint_last.pth) contains the model state from the final training step. You can disable EMA weights by setting use_ema=False, in which case the best checkpoint will use standard model weights instead.
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 →