This feature implements the BEMA algorithm to update the reference model during DPO training.
from trl.experimental.bema_for_ref_model import BEMACallback, DPOTrainer
from datasets import load_dataset
dataset = load_dataset("trl-internal-testing/zen", "standard_preference", split="train")
bema_callback = BEMACallback(update_ref_model=True)
trainer = DPOTrainer(
model="trl-internal-testing/tiny-Qwen2ForCausalLM-2.5",
train_dataset=dataset,
callbacks=[bema_callback],
)
trainer.train()( resume_from_checkpoint: str | bool | None = None trial: optuna.Trial | dict[str, Any] | None = None ignore_keys_for_eval: list[str] | None = None ) → ~trainer_utils.TrainOutput
Parameters
str or bool, optional) —
If a str, local path to a saved checkpoint as saved by a previous instance of Trainer. If a
bool and equals True, load the last checkpoint in args.output_dir as saved by a previous instance
of Trainer. If present, training will resume from the model/optimizer/scheduler states loaded here. optuna.Trial or dict[str, Any], optional) —
The trial run or the hyperparameter dictionary for hyperparameter search. list[str], optional) —
A list of keys in the output of your model (if it is a dictionary) that should be ignored when
gathering predictions for evaluation during the training. Returns
~trainer_utils.TrainOutput
Object containing the global step count, training loss, and metrics.
Main training entry point.
Will save the model, so you can reload it using from_pretrained().
Will only save from the main process.
( commit_message: str | None = 'End of training' blocking: bool = True token: str | None = None revision: str | None = None **kwargs )
Parameters
str, optional, defaults to "End of training") —
Message to commit while pushing. bool, optional, defaults to True) —
Whether the function should return only when the git push has finished. str, optional, defaults to None) —
Token with write permission to overwrite Trainer’s original args. str, optional) —
The git revision to commit from. Defaults to the head of the “main” branch. dict[str, Any], optional) —
Additional keyword arguments passed along to ~Trainer.create_model_card. Upload self.model and self.processing_class to the 🤗 model hub on the repo self.args.hub_model_id.
( update_freq: int = 400 ema_power: float = 0.5 bias_power: float = 0.2 lag: int = 10 update_after: int = 0 multiplier: float = 1.0 min_ema_multiplier: float = 0.0 device: str = 'cpu' update_ref_model: bool = False ref_model_update_freq: int = 400 ref_model_update_after: int = 0 )
Parameters
int, optional, defaults to 400) —
Update the BEMA weights every X steps. Denoted this as {@html "ϕ"} in the paper. float, optional, defaults to 0.5) —
Power for the EMA decay factor. Denoted {@html "κ"} in the paper. To disable EMA, set this to 0.0. float, optional, defaults to 0.2) —
Power for the BEMA scaling factor. Denoted {@html "η"} in the paper. To disable BEMA, set this to 0.0. int, optional, defaults to 10) —
Initial offset in the weight decay schedule that controls early-stage smoothness by acting as a virtual
starting age for the updates. Denoted as {@html "ρ"} in the paper. int, optional, defaults to 0) —
Burn-in time before starting to update the BEMA weights. Denoted {@html "τ"} in the paper. float, optional, defaults to 1.0) —
Initial value for the EMA decay factor. Denoted as {@html "γ"} in the paper. float, optional, defaults to 0.0) —
Minimum value for the EMA decay factor. str, optional, defaults to "cpu") —
Device to use for the BEMA buffers, e.g. "cpu" or "cuda". Note that in most cases, this device SHOULD
BE DIFFERENT from the device used for training in order to avoid OOM. bool, optional, defaults to False) —
Whether to update the reference model with BEMA weights. This creates a lagged, smoothed version of the
main model as the reference model. int, optional, defaults to 400) —
Update the reference model with BEMA weights every this many steps. int, optional, defaults to 0) —
Number of steps to wait before starting to update the reference model. A TrainerCallback that implements BEMA (Bias-Corrected Exponential Moving Average) by Adam Block and Cyril Zhang. Code from https://github.com/abblock/bema under MIT license.
BEMA computes model weights that scale like:
where is the current model weights, is a snapshot of the model weights at the
first update_after step, is the exponential moving average of the model weights, and is a scaling factor that decays with the number of steps as
The EMA is computed as:
where is a decay factor that decays with the number of steps as