How to Enable EMA (Exponential Moving Average) During RF-DETR Training
RF-DETR enables EMA automatically via TrainConfig.use_ema=True (the default), which injects the RFDETREMACallback into the Lightning trainer to maintain a shadow averaged model alongside the main weights.
Exponential Moving Average (EMA) stabilizes training and improves final model quality by maintaining a slowly updated copy of the network parameters. In the roboflow/rf-detr repository, EMA is implemented as a first-class citizen through PyTorch Lightning callbacks, requiring zero manual intervention when using the standard training API.
Configuring EMA Behavior
RF-DETR exposes EMA controls through the TrainConfig class, allowing you to toggle the feature and fine-tune its aggression.
The use_ema Toggle
EMA is controlled by the boolean field use_ema defined in src/rfdetr/config.py (lines 969–971). By default, this value is True, meaning EMA is active for all training runs unless explicitly disabled.
from rfdetr.config import TrainConfig
# EMA is enabled by default
config = TrainConfig(dataset_dir="path/to/data", output_dir="output")
assert config.use_ema is True
EMA Hyperparameters
When EMA is active, three parameters govern its behavior, also located in src/rfdetr/config.py (around lines 1010–1018):
ema_decay(float, default0.993): The base decay factor determining how much weight the shadow model retains versus new parameters.ema_tau(int, default100): The warm-up time constant (τ) for ramping the decay rate from a small value to the base decay over the first τ steps.ema_update_interval(int, default1): The number of optimizer steps between EMA updates; increasing this reduces computational overhead.
Internal Implementation Details
The EMA mechanism is encapsulated in RFDETREMACallback, found in src/rfdetr/training/callbacks/ema.py. This callback wraps torch.optim.swa_utils.AveragedModel and implements a custom averaging function via the _avg_fn method and decay scheduling logic in _effective_decay.
The trainer construction logic in build_trainer (located in src/rfdetr/training/__init__.py) automatically appends this callback to the Lightning trainer when train_config.use_ema is True. Notably, the implementation includes a guard that disables EMA when using sharded training strategies (e.g., FSDP or DeepSpeed), as confirmed by the unit test in tests/training/callbacks/test_ema_callback.py (lines 306–312).
Practical Training Examples
Using the High-Level Train API
For most users, enabling and configuring EMA requires only passing arguments to TrainConfig:
from rfdetr import train
from rfdetr.config import TrainConfig, RFDETRSmallConfig
tc = TrainConfig(
dataset_dir="path/to/dataset",
output_dir="output",
use_ema=True, # Optional: True is the default
ema_decay=0.995, # Stronger averaging
ema_tau=200, # Slower warm-up
ema_update_interval=2, # Update every 2 steps
)
mc = RFDETRSmallConfig()
train(mc, tc) # EMA callback is injected automatically
Explicit Trainer Construction
If you are manually building the Lightning trainer, EMA is still handled automatically, though you can inspect or modify the callback stack:
from rfdetr.training import build_trainer
from rfdetr.training.module_model import RFDETRModelModule
from rfdetr.config import TrainConfig, RFDETRSmallConfig
tc = TrainConfig(dataset_dir="data", use_ema=True)
mc = RFDETRSmallConfig()
model = RFDETRModelModule(model_config=mc, train_config=tc)
trainer = build_trainer(model, tc) # Callback injected here
trainer.fit(model)
Disabling EMA for Faster Iteration
To disable EMA—for debugging, rapid prototyping, or when memory is constrained—set use_ema=False:
from rfdetr import train
from rfdetr.config import TrainConfig, RFDETRSmallConfig
tc = TrainConfig(dataset_dir="data", use_ema=False) # EMA skipped
train(RFDETRSmallConfig(), tc)
Summary
- Default Enabled: EMA is active by default in RF-DETR through
TrainConfig.use_ema=True. - Configurable Decay: Tune
ema_decay,ema_tau, andema_update_intervalinsrc/rfdetr/config.pyto control averaging strength and frequency. - Automatic Injection: The
build_trainerhelper automatically addsRFDETREMACallbackfromsrc/rfdetr/training/callbacks/ema.pywhen EMA is enabled. - Sharded Training Guard: EMA is automatically disabled for sharded strategies to prevent synchronization issues.
- Zero Overhead Setup: No manual callback registration is required; the framework handles weight swapping for validation and checkpointing.
Frequently Asked Questions
Is EMA enabled by default in RF-DETR?
Yes. The TrainConfig class initializes use_ema=True by default (see src/rfdetr/config.py lines 969–971). You must explicitly set it to False to disable the feature.
What is the difference between ema_decay and ema_tau?
ema_decay is the asymptotic decay factor (default 0.993) that controls the long-term mixing ratio between the shadow model and live parameters. ema_tau (default 100) defines a warm-up period during which the effective decay is linearly ramped from a small value up to the full ema_decay, preventing early-training instability.
Why does EMA get disabled with sharded training strategies?
The build_trainer function checks for sharded strategies (FSDP, DeepSpeed) and skips EMA callback injection because maintaining a separate averaged model across distributed shards requires complex state synchronization that is not implemented in the current callback logic, as verified in tests/training/callbacks/test_ema_callback.py.
How do I know if the EMA model is being used for validation?
When EMA is active, RFDETREMACallback automatically swaps the EMA weights into the model during the validation and test phases. Metrics logged during validation (e.g., val/ema_mAP) reflect the EMA model performance, and the final checkpoint saved by the trainer contains the EMA weights, not the raw optimizer weights.
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 →