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 and RFDETRModelModule (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:

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 →