Direct Preference Optimization (DPO), introduced by Rafailov, Sharma, Mitchell, Ermon, Manning and Finn in 2023, replaced the reward model and reinforcement learning loop of classic RLHF with a single classification-style loss over preference pairs. That made preference tuning look like supervised fine-tuning, and many teams now run it on the same GPUs and training stack they use for SFT.
It is not quite SFT, though. Every step needs log-probabilities from two models on two completions per prompt, the loss depends on whole-sequence sums that are sensitive to precision, and the vocabulary-sized logits tensor is usually the largest thing in memory. This article covers DPO as a GPU workload: what a step computes, where the memory and FLOPs go, how to host the reference model, how to shard it across GPUs, and what fails in practice. The mathematics, including the derivation of the loss and its variants, is covered in DPO math; here we only restate what the hardware has to compute.
What one DPO step computes
A preference example is a prompt x, a preferred completion y_w and a rejected completion y_l. For each completion the step needs the sum of token log-probabilities under the policy being trained, log pi(y|x), and under a frozen reference policy, log pi_ref(y|x), usually the SFT checkpoint you started from. The loss is:
loss = -log sigmoid( beta * [ (log pi(y_w|x) - log pi_ref(y_w|x)) - (log pi(y_l|x) - log pi_ref(y_l|x)) ] )The bracketed quantity is the margin between two implicit rewards, each defined as beta times the log-ratio of policy to reference. So every example needs four sequence scores: policy on chosen, policy on rejected, reference on chosen, reference on rejected. Only the two policy forwards need gradients; the reference forwards are pure inference.
Three details turn that formula into GPU work. The sums run only over completion tokens, but the prompt must still be in the forward pass because the completion is conditioned on it. Implementations usually concatenate chosen and rejected sequences into one batch of 2B rows, so one forward and backward serve both halves. And DPOConfig disables dropout by default, because dropout would perturb policy and reference scores differently and corrupt the log-ratio.
A worked example of the numbers
Take beta = 0.1 and one pair. Summed over completion tokens, the policy scores the chosen answer at -42.0 and the rejected answer at -47.5; the reference scores them at -43.0 and -46.0. The implicit rewards are 0.1 x (-42.0 + 43.0) = 0.10 for chosen and 0.1 x (-47.5 + 46.0) = -0.15 for rejected, so the margin is 0.25 and the loss is -log sigmoid(0.25) = 0.576. At the start of training the policy equals the reference, every margin is zero, and the loss is ln 2 = 0.693, which is the first number you should see in your logs.
The gradient on this pair is scaled by sigmoid(-margin) = 0.438. Pairs the model already ranks correctly by a wide margin contribute little, and pairs it gets wrong contribute close to the full weight. This self-weighting is why DPO is stable, but it also means a batch of easy pairs does almost nothing while costing the full compute of four forwards and a backward.
For the GPU this means two things. The loss is a difference of differences between sums of tens to hundreds, so accumulate them in fp32 even when the model runs in bf16, whose 8-bit mantissa loses the sub-unit changes that carry the signal. And the per-token log-softmax spans the whole vocabulary, which is where the memory goes.
Where the memory goes for an 8B model
Use Llama 3.1 8B as a concrete case: about 8.03 billion parameters and a 128,256-token vocabulary. Full-parameter DPO with AdamW in mixed precision stores roughly 16 bytes per trainable parameter (bf16 weights and gradients, fp32 master weights, and two fp32 Adam moments), which is about 120 GiB for the policy alone. A separate bf16 reference model adds about 15 GiB. That is far beyond a single 80 GB GPU, so full DPO at this size is a multi-GPU job with sharding, exactly as full SFT is; see ZeRO, in depth for the per-stage arithmetic.
The part people underestimate is the logits. For one 2,048-token sequence they are 2,048 x 128,256 values: 0.49 GiB in bf16 and 0.98 GiB in fp32. A micro-batch of four pairs is eight sequences, or 7.8 GiB of fp32 logits, before the log-softmax output and its gradient. Gradient checkpointing, on by default in DPOConfig, shrinks layer activations but not the output projection, so at long context the logits can exceed all layer activations combined.
| Item (Llama 3.1 8B, bf16 compute) | Size | Can you remove it? |
|---|---|---|
| Policy weights, grads, fp32 master, Adam moments (full DPO) | about 120 GiB | Shard with FSDP or ZeRO-3, or train LoRA adapters instead |
| Frozen base weights (LoRA DPO) | about 15 GiB | Quantize to 4-bit (QLoRA-style) at some quality and speed cost |
| Separate reference model, bf16 | about 15 GiB | Yes: adapter-disable trick or precomputed log-probs |
| fp32 logits, 8 sequences of 2,048 tokens | about 7.8 GiB | Chunk the log-softmax or use a fused kernel |
| Layer activations with checkpointing | grows with batch and length | Smaller micro-batches, offloading |
Four ways to host the reference model
The reference model is what makes DPO heavier than SFT, and there are four common ways to pay for it.
- A second resident copy. Load the SFT checkpoint twice, one trainable and one frozen in eval mode. Simple and flexible, but it costs a full set of weights per GPU, or a second sharded model under FSDP. In TRL, leaving
ref_model=Nonefor a full model makes the trainer use a copy of the initial policy. - The same weights with the adapter switched off. When you train LoRA adapters, the frozen base plus a disabled adapter is exactly the SFT model, so the reference forward is the same network with the adapter bypassed. The reference then costs no extra memory, only compute. TRL does this for PEFT models without a standalone reference, which is also why
sync_ref_model=Trueis not supported in that setup: there is no separate reference to update. - Precomputed log-probabilities. Because the reference is frozen, its four numbers per pair never change. Compute them once in a no-grad pass over the dataset, store two floats per example, and drop the reference from the training loop.
precompute_ref_log_probs=Truedoes this, withprecompute_ref_batch_sizeallowed to be larger than the training batch because the pass holds no activations for backward. It is not supported with streamingIterableDatasetinputs, and it is incompatible withsync_ref_model, which moves the reference during training. - A remote scorer. The reference runs on an inference server, the pattern used in RLHF pipelines on GPU. The server must match the trainer's tokenizer, chat template and numerics, or every log-ratio picks up a bias.
On FLOPs: per token, a training step on the policy costs roughly 6P FLOPs (forward plus backward) plus about 2P more when checkpointing recomputes the forward, and a reference forward costs about 2P. So the reference is roughly 20 to 25 percent of step compute. Precomputing does not remove that cost for a single epoch, it moves it to the start; it pays off in memory, in simpler sharding and in every additional epoch or hyperparameter sweep over the same data.
A training configuration that fits
This uses TRL's DPOTrainer as documented on its main branch in late 2026; names change between releases, so check yours. It trains LoRA adapters on an 8B model with a precomputed reference:
import torch
from datasets import load_dataset
from peft import LoraConfig
from trl import DPOConfig, DPOTrainer
args = DPOConfig(
output_dir="dpo-8b-lora",
model_init_kwargs={"dtype": torch.bfloat16}, # a string model id loads in float32 otherwise
beta=0.1, # default; larger stays closer to the reference
learning_rate=1e-5, # DPOConfig default is 1e-6; adapters need more
per_device_train_batch_size=2, # pairs per GPU; each pair is two sequences
gradient_accumulation_steps=16,
max_length=1024, # default; truncation happens before padding
gradient_checkpointing=True, # default in DPOConfig
precompute_ref_log_probs=True, # one reference pass up front, no reference in the loop
precompute_ref_batch_size=16, # the no-grad pass can use a much larger batch
logging_steps=10,
)
trainer = DPOTrainer(
model="meta-llama/Llama-3.1-8B-Instruct", # your SFT checkpoint in practice
args=args,
train_dataset=load_dataset("trl-lib/ultrafeedback_binarized", split="train"),
peft_config=LoraConfig(r=16, lora_alpha=32, target_modules="all-linear"),
)
trainer.train()Two defaults deserve attention. DPOConfig's learning_rate defaults to 1e-6, which suits full-parameter DPO; the TRL docs note that adapters usually need a higher rate, around 1e-5. And if you pass the model as a string without setting dtype in model_init_kwargs, the trainer loads it in float32, which doubles weight memory and is the most common reason a run that should fit does not.
Writing the loss by hand once makes the memory behaviour obvious. This version chunks positions so the full fp32 log-softmax never exists at once, and sums in fp32:
import torch
import torch.nn.functional as F
from torch.utils.checkpoint import checkpoint
def _chunk_logps(lg, lbl, m):
lp = torch.gather(F.log_softmax(lg.float(), dim=-1), 2, lbl.unsqueeze(-1)).squeeze(-1)
return (lp * m).sum(dim=1) # fp32 partial sums
def sequence_logps(logits, labels, mask, chunk=256):
"""Sum of log p(token) over completion tokens only. Each chunk is checkpointed, so its
fp32 log-softmax is freed after forward and recomputed in backward."""
logits = logits[:, :-1] # position t predicts token t+1
labels = labels[:, 1:]
mask = mask[:, 1:].float()
total = torch.zeros(logits.size(0), device=logits.device, dtype=torch.float32)
for s in range(0, logits.size(1), chunk):
total = total + checkpoint(_chunk_logps, logits[:, s:s + chunk],
labels[:, s:s + chunk], mask[:, s:s + chunk],
use_reentrant=False)
return total
def dpo_step(policy, batch, ref_chosen, ref_rejected, beta=0.1):
# batch holds chosen rows first, then rejected rows: one forward for both halves
out = policy(input_ids=batch["input_ids"], attention_mask=batch["attention_mask"])
logps = sequence_logps(out.logits, batch["input_ids"], batch["completion_mask"])
pi_c, pi_r = logps.chunk(2)
margin = beta * ((pi_c - ref_chosen) - (pi_r - ref_rejected))
loss = -F.logsigmoid(margin).mean()
loss.backward()
return {"loss": loss.item(),
"reward_acc": (margin > 0).float().mean().item(),
"logps_chosen": pi_c.mean().item()} # watch this: it should not collapseWith use_liger_kernel=True, TRL uses Liger's chunked loss path, which never materializes the full logits, at the cost of a few logged metrics and options.
Scaling out: sharding, sequence length and batch shape
For full-parameter DPO beyond a few billion parameters, shard the policy with FSDP or ZeRO-3 as for SFT; PyTorch FSDP, in depth covers the mechanics. A resident reference should be sharded too, for inference only and outside the optimizer.
Sequence length hits DPO twice: every pair is two sequences, and logits grow linearly with length. FlashAttention-style kernels handle attention memory but not the vocabulary projection, so chunked or fused loss matters more as context grows. Truncation is a trap too: max_length defaults to 1,024 and keeps the start, so if two answers differ only near their ends, the pair carries no signal but still costs a full step.
Keep the chosen and rejected rows of a pair in the same forward so padding and numerics match, and bucket by length to cut padding waste. TRL currently has padding_free disabled (it warns and falls back to padding), so do not plan capacity around it.
Failure modes you will see on real runs
- Both log-probs fall. Reward accuracy climbs while
logps/chosensteadily decreases, because the loss only cares about the margin and pushing the rejected answer down is easier. Watch the chosen log-prob, lower the learning rate or raise beta, or add a small SFT term with the multi-loss option (loss_type=["sigmoid", "sft"]). - Length exploitation. Sequence sums grow with length, so if chosen answers are systematically longer, the model learns verbosity. Check the length distribution of chosen vs rejected before training; length-normalized variants such as
sigmoid_normexist for this. - Loss stuck at 0.693. The margin never moves: the adapter is not in the optimizer, or the learning rate is too small for LoRA.
- Reference mismatch. Scoring the reference with a different tokenizer, chat template, dtype or attention implementation than the policy shifts every log-ratio. The first logged loss should be almost exactly ln 2; if it is not, the two paths disagree.
- Out-of-memory at the loss, not in the layers. The stack trace ends in log_softmax or cross-entropy. Chunk and checkpoint the loss, use the fused kernel, or cut the micro-batch; layer checkpointing will not help.
- Noisy labels. If annotators often disagree, consider
loss_type="robust"withlabel_smoothingset to your estimated flip rate.
Operational guidance and trade-offs
Gate DPO runs on evals. Log TRL's rewards/margins, rewards/accuracies and logps/chosen, and add a held-out preference set plus a small capability suite, because reward accuracy can reach 90 percent while general ability degrades. Checkpoint often; runs frequently peak early and then over-optimize.
LoRA with an adapter-disabled or precomputed reference fits an 8B model on one GPU and is usually enough for style, format and safety preferences; full-parameter DPO shifts behaviour further but needs sharding and tighter learning-rate control. Plain DPO trains on fixed pairs, which is cheap, but they drift from what the current policy generates, so many teams regenerate pairs each round, paying generation cost that brings back part of the RLHF pipeline DPO replaced.
What to do next
- Start from your SFT checkpoint and a preference set in prompt, chosen and rejected format with an explicit prompt.
- Check chosen vs rejected length statistics and drop near-duplicate pairs before you spend GPU time.
- Run one step and confirm the first logged loss is about 0.693 and the reward accuracy is about 0.5.
- Pick a reference strategy: adapter-disable for LoRA, precomputed log-probs for multi-epoch runs, a sharded resident copy only if you need a moving reference.
- Set
model_init_kwargsdtype explicitly, keep log-prob sums in fp32 and chunk or fuse the loss before raising the context length. - Track
logps/chosenalongside margins, and stop when held-out preference accuracy or your capability suite stops improving.