Direct Preference Optimization (DPO) trains a language model on pairs of responses, one preferred and one rejected, without a separate reinforcement learning loop. In its usual form the pairs are collected once, before training, often from a different model. Online DPO removes that limitation. At every training step it samples two fresh responses from the current policy, asks an annotator which is better, and takes one DPO step on those pairs. The method was introduced as online AI feedback (OAIF) by Guo et al. in 2024 (arXiv 2402.04792). The paper reports that it beat both offline DPO-style methods and RLHF in human evaluations, and that the annotator's behaviour can be steered just by changing its instruction prompt.
This article is about running it on GPUs: the loss, a minimal step, annotators, the TRL trainer and vLLM, the memory budget and weight sync. One correction first. Online DPO in this sense does not mean learning from live user feedback on a deployed model. The pairs come from the policy being trained, and an annotator model labels them inside the training loop.
Offline, iterative and online
Three regimes are easy to confuse, and they need different infrastructure.
| Regime | Where pairs come from | When they are labelled | Infrastructure |
|---|---|---|---|
| Offline DPO | Fixed dataset, often another model or humans | Before training | Trainer only |
| Iterative DPO | Current policy, once per round | Between rounds, in bulk | Trainer, plus a batch generation job per round |
| Online DPO | Current policy, every step | Inside the step | Trainer, generator and annotator running together |
The reason to go online is distribution shift. DPO only adjusts the probabilities of the responses in its pairs, so if they came from another model, the policy's own mistakes may never appear in a pair. Online sampling keeps every pair on-policy, at the price of generation and annotation inside every step. Iterative DPO is the middle ground: on-policy data refreshed once per round, with much simpler scheduling.
The objective and a worked loss
For a prompt x with chosen response y_w and rejected response y_l, DPO compares how much the policy has raised each response's log-probability relative to a frozen reference model:
margin = beta * [ (log pi(y_w|x) - log ref(y_w|x)) - (log pi(y_l|x) - log ref(y_l|x)) ]
loss = -log sigmoid(margin)Here log pi(y|x) is the sum of token log-probabilities of the response. Take beta = 0.1, a policy that gives the chosen response -42.0 against the reference's -43.5, and gives the rejected response -40.0 against the reference's -39.0. The margin is 0.1 * (1.5 - (-1.0)) = 0.25, and the loss is log(1 + e-0.25) = 0.576. The gradient pushes the chosen response up and the rejected one down, weighted by sigmoid(-margin) = 0.44, so pairs the model already orders correctly by a wide margin contribute little.
Beta controls how far the policy may drift from the reference; a larger beta keeps it closer. Online, KL to the reference is a primary health metric. The IPO loss is an alternative when pairs become nearly separable and the sigmoid loss keeps pushing margins up.
One step as a data flow
Each step samples prompts, generates two completions per prompt at a temperature high enough to make them differ (TRL defaults to 0.9), scores both, orders them (dropping near-ties), takes a DPO step, and hands the new weights to the generator.
The stages are sequential. The generator cannot start the next batch until it has the new weights, unless you accept one step of staleness, which makes GPUs harder to keep busy than in offline training.
A minimal training step
A minimal step in plain PyTorch makes each cost explicit. generator and annotator are interfaces you supply. The point is the order of operations and what runs with gradients.
import torch
import torch.nn.functional as F
def completion_logp(model, ids, attn, comp_mask):
logits = model(input_ids=ids, attention_mask=attn).logits[:, :-1]
# gather the target logit and subtract logsumexp one sequence at a time:
# peak fp32 is T x V instead of the B x T x V of a full log_softmax
tok = logits.gather(-1, ids[:, 1:].unsqueeze(-1)).squeeze(-1).float()
lse = torch.stack([torch.logsumexp(row.float(), -1) for row in logits])
return ((tok - lse) * comp_mask[:, 1:]).sum(-1)
def online_dpo_step(policy, ref, generator, annotator, tok, prompts, opt,
beta=0.1, min_gap=0.0):
a, b = generator.sample(prompts, n=2, temperature=0.9, max_new_tokens=256)
sa, sb = annotator.score(prompts, a), annotator.score(prompts, b) # tensors [B]
keep = (sa - sb).abs() > min_gap # drop ties
if keep.sum() == 0:
return None
win_a = sa >= sb
chosen = [x if w else y for x, y, w in zip(a, b, win_a)]
rejected = [y if w else x for x, y, w in zip(a, b, win_a)]
batch = tok.pack(prompts, chosen, rejected, keep) # ids, attn, comp_mask x 2
pc = completion_logp(policy, *batch.chosen)
pr = completion_logp(policy, *batch.rejected)
with torch.no_grad():
rc = completion_logp(ref, *batch.chosen)
rr = completion_logp(ref, *batch.rejected)
margin = beta * ((pc - rc) - (pr - rr))
loss = -F.logsigmoid(margin).mean()
loss.backward()
torch.nn.utils.clip_grad_norm_(policy.parameters(), 1.0)
opt.step(); opt.zero_grad(set_to_none=True)
generator.load_weights(policy.state_dict()) # the online-only step
return {"loss": loss.item(), "acc": (margin > 0).float().mean().item(),
"kept": keep.float().mean().item()}Count the passes per step for B prompts. There are 2B sampled sequences, 2B annotator forwards, 2B policy forward and backward passes, and 2B reference forwards. Compare this with the four forward passes and two backward passes of offline DPO training. Online DPO adds token-by-token generation, which is memory-bandwidth bound, plus a whole second model's forward pass.
Choosing the annotator
Reward model. A sequence classifier that returns one scalar per response. It is fast and batchable, and comparing two scalars gives a preference directly. Its weakness is that the policy is optimised against it continuously, so any exploitable quirk gets found. See reward model training for how those models are built. If the reward model uses a different tokenizer or chat template from the policy, completions must be decoded to text and re-encoded for it.
LLM judge. The OAIF setting: a prompted model sees both responses and picks one. It needs no reward-model training, but every judgement is a generation call. Judges often favour the first-listed and the longer response, so score both orders, keep agreeing verdicts and watch length.
Rules and penalties. Cheap programmatic checks, such as format validity, a missing end-of-sequence token or a banned phrase, can be added to a model score. A response that hits the length cap without an EOS token is usually worse than it looks to a scorer. A fixed penalty for it is one of the most effective guards against length blow-up.
Running it with TRL
Hugging Face TRL implements this as OnlineDPOTrainer. In current TRL documentation it lives in trl.experimental.online_dpo, takes a prompt-only dataset, and receives its annotator through reward_funcs. That argument accepts a reward model, a model path, or Python callables, which are summed using reward_weights. Older TRL releases exposed the trainer at the package top level with different annotator arguments, so check the documentation for the version you have installed before copying code.
from datasets import load_dataset
from transformers import AutoModelForSequenceClassification, AutoTokenizer
from trl.experimental.online_dpo import OnlineDPOConfig, OnlineDPOTrainer
tok = AutoTokenizer.from_pretrained("my-org/policy-sft")
rm = AutoModelForSequenceClassification.from_pretrained("my-org/rm", num_labels=1)
prompts = load_dataset("my-org/prompts", split="train") # column: "prompt"
args = OnlineDPOConfig(
output_dir="policy-online-dpo",
beta=0.1, loss_type="sigmoid", learning_rate=5e-7,
max_new_tokens=256, temperature=0.9, missing_eos_penalty=1.0,
per_device_train_batch_size=8, gradient_accumulation_steps=2,
use_vllm=True, vllm_mode="colocate",
vllm_gpu_memory_utilization=0.30, vllm_enable_sleep_mode=True,
logging_steps=10, report_to="wandb",
)
trainer = OnlineDPOTrainer(model="my-org/policy-sft", reward_funcs=rm, args=args,
processing_class=tok, train_dataset=prompts)
trainer.train()Defaults worth knowing from the current config are beta=0.1, temperature=0.9, max_new_tokens=64 (too short for most chat tasks), learning_rate=5e-7, gradient checkpointing on, and vllm_gpu_memory_utilization=0.55 when colocated. Launch with accelerate launch for multi-GPU runs.
Where GPU memory and time go
Work through the memory for a 7B policy on one node of eight 80 GB GPUs with ZeRO-3. Mixed-precision AdamW needs about 16 bytes per parameter for weights, gradients, fp32 master weights and the two moments. That is 112 GB, or 14 GB per GPU once sharded. A bf16 reference model is 14 GB, about 2 GB per GPU if sharded the same way. A 7B reward model kept whole on each GPU for fast scoring is another 14 GB. Colocated vLLM at the default 0.55 utilization claims 44 GB for its own copy of the weights and its KV cache. That totals about 74 GB before activations and logits. Logits alone for 16 sequences of 512 tokens over a 150k vocabulary are about 2.5 GB in bf16, before any fp32 upcast. The default does not fit here. The levers, in order of cost:
- Lower
vllm_gpu_memory_utilizationto 0.25-0.35. The KV cache shrinks, so fewer sequences decode concurrently, and generation slows. - Use a smaller reward model, or a judge served elsewhere.
- Enable
vllm_enable_sleep_mode, which offloads vLLM's weights and cache during the optimizer step at the cost of host-device transfers to wake it. - Switch to
vllm_mode='server': dedicate some GPUs to a vLLM server and train on the rest. Weights are pushed to the server over a communication group each step.
Time follows the same logic. Decoding is bandwidth-bound and often the largest share of the step, so throughput depends on how many sequences decode in parallel, which KV-cache space sets. Under ZeRO-3, ds3_gather_for_generation=True gathers full weights for generation. Turning it off fits larger models but is much slower and incompatible with vLLM. Profile generate, annotate, train and sync separately before tuning.
Weight sync: the online-only failure
The generator is a separate engine with its own copy of the weights, and it must receive every update. If a sharded parameter is skipped, a LoRA adapter never merged, or a server push fails silently, training continues with plausible metrics on off-policy data, and the benefit of going online quietly disappears.
Test it directly. Take a few completions the generator just produced, recompute their token log-probabilities with the trainer's policy, and compare with the log-probabilities the generator reported, with both computed on the same distribution (raw logits, or the same temperature on both sides). per_token_logp is completion_logp without the final sum. Small differences from kernels and precision are normal. A gap that grows with step count means the generator is serving stale weights.
def check_sync(trainer_policy, gen_out, tol=0.05):
# gen_out: token ids, completion mask and per-token logprobs reported by vLLM
lp_train = per_token_logp(trainer_policy, gen_out.ids, gen_out.mask)
diff = (lp_train - gen_out.logprobs).abs()[gen_out.mask.bool()].mean().item()
if diff > tol:
raise RuntimeError(f"generator weights look stale: mean |dlogp| = {diff:.3f}")
return diff
Metrics to watch
TRL logs the signals you need. Read them together, not alone.
objective/scoresandobjective/scores_margin: what the annotator thinks. Rising scores are necessary but not sufficient.objective/kl: divergence from the reference. Steady growth is expected. A sudden jump usually means reward exploitation or a learning rate that is too high.rewards/accuraciesandrewards/margins: DPO's implicit reward. If accuracy sits near 1.0 early, pairs have become trivially easy.objective/entropy: falling entropy means the two samples converge. Identical pairs carry no signal, so raise the temperature or the sample count.val/contain_eos_token: if this drops, completions are hitting the length cap.
Add periodic evaluation by a different judge or humans on held-out prompts. If the training score rises while that stays flat, the policy is exploiting the annotator.
Failure modes
- Annotator exploitation. Longer, more confident or more flattering responses score higher without being better. Defend with an independent evaluation set, a length penalty or length-controlled comparison, and a higher beta.
- Stale generator weights. Covered above. Assert sync every N steps.
- Collapsed diversity. Two samples from a sharp policy are often the same text. Track the share of pairs dropped as ties and the entropy.
- Template mismatch. The policy, the annotator and the reference must see the same chat formatting their training assumed. A reward model scoring raw unformatted text gives noisy preferences.
- Out-of-memory crashes mid-run. Long prompts or completions late in the dataset overflow a budget that fit early. Cap
max_lengthand test with the longest prompts first.
Trade-offs
| Choice | Gain | Cost |
|---|---|---|
| Online vs iterative | Fully on-policy pairs every step | Generation inside the step; sync complexity |
| Reward model vs LLM judge | Speed, batchability | Exploitable; needs training data |
| Colocate vs server vLLM | No idle GPUs, simpler launch | Memory contention; slower decoding |
| Sigmoid vs IPO loss | Standard, well understood | IPO resists overfitting on separable pairs but needs its own tau tuning |
| Online DPO vs GRPO or PPO | No value model, pairwise signal | Uses only 2 samples per prompt; less signal per prompt than group methods |
If your annotator is a verifier with exact answers, such as math or code tests, compare with GRPO, which uses groups of samples per prompt. If you need the full classical stack with a value model, see the RLHF pipeline.
What to do next
- Run offline DPO on your preference data first and keep its evaluation as the baseline.
- Choose the annotator. For a judge, score both orders and keep agreeing verdicts. For a reward model, confirm it uses the policy's chat template.
- Write the memory budget per GPU, as above, before launching. Choose colocate or server mode from the result.
- Set max_new_tokens to your real response length, add missing_eos_penalty, and start at beta 0.1.
- Add the weight-sync check and fail the run if it trips.
- Track scores, KL, entropy, tie rate and EOS rate together, plus an independent judge on held-out prompts every few hundred steps.
- Stop or raise beta when the training score rises while the independent evaluation stalls.