# How to Train RF-DETR with a Custom Dataset: A Complete Guide

> Learn how to train RF-DETR with your own data. Follow our guide for a quick one-line start or gain full control with PyTorch Lightning for custom dataset training.

- Repository: [Roboflow/rf-detr](https://github.com/roboflow/rf-detr)
- Tags: how-to-guide
- Published: 2026-09-08

---

**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.

```python
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`](https://github.com/roboflow/rf-detr/blob/main/src/rfdetr/training/module_data.py) automatically detects the structure and builds the appropriate PyTorch datasets.

### COCO Format

Place [`train/_annotations.coco.json`](https://github.com/roboflow/rf-detr/blob/main/train/_annotations.coco.json) inside your dataset root. Include [`valid/_annotations.coco.json`](https://github.com/roboflow/rf-detr/blob/main/valid/_annotations.coco.json) for validation splits.

### YOLO Format

Include a [`data.yaml`](https://github.com/roboflow/rf-detr/blob/main/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`](https://github.com/roboflow/rf-detr/blob/main/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`](https://github.com/roboflow/rf-detr/blob/main/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`](https://github.com/roboflow/rf-detr/blob/main/src/rfdetr/training/module_model.py))

Located in [`src/rfdetr/training/module_model.py`](https://github.com/roboflow/rf-detr/blob/main/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:

```python
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:

```python
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:

```python
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.

```python
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`](https://github.com/roboflow/rf-detr/blob/main/module_data.py) and [`module_model.py`](https://github.com/roboflow/rf-detr/blob/main/module_model.py).

## Monitoring and Checkpoints

Training metrics are visualized using helpers in [`src/rfdetr/visualize/training.py`](https://github.com/roboflow/rf-detr/blob/main/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:

```python
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`](https://github.com/roboflow/rf-detr/blob/main/src/rfdetr/training/module_data.py)) for data loading and `RFDETRModelModule` ([`src/rfdetr/training/module_model.py`](https://github.com/roboflow/rf-detr/blob/main/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`](https://github.com/roboflow/rf-detr/blob/main/_annotations.coco.json) files) and **YOLO** format (with [`data.yaml`](https://github.com/roboflow/rf-detr/blob/main/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.