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.

Advertisement

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.

One DPO step: two models, four sequence forwards, one scalar lossPreference batchprompt, chosen, rejectedConcatenate + pad2B sequences, left padPolicy forwardgrad on, checkpointedReference forwardno_grad, or cachedMasked log-softmaxsum over completionRef log-probsfp32, per sequenceDPO loss-log sigmoid(beta * margin)Backward + optimizerpolicy params onlypolicyreferencelogitslogitspi logpsref logpsscalarThe reference path holds no activations for backward; the policy path holds all of them.Logits of shape [2B, T, V] are the largest single tensor in the step unless you chunk them.
A DPO step as a data flow. The reference branch is inference only, so it can be cached ahead of time; the policy branch holds activations for backward.

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.

Advertisement

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)SizeCan you remove it?
Policy weights, grads, fp32 master, Adam moments (full DPO)about 120 GiBShard with FSDP or ZeRO-3, or train LoRA adapters instead
Frozen base weights (LoRA DPO)about 15 GiBQuantize to 4-bit (QLoRA-style) at some quality and speed cost
Separate reference model, bf16about 15 GiBYes: adapter-disable trick or precomputed log-probs
fp32 logits, 8 sequences of 2,048 tokensabout 7.8 GiBChunk the log-softmax or use a fused kernel
Layer activations with checkpointinggrows with batch and lengthSmaller 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.

  1. 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=None for a full model makes the trainer use a copy of the initial policy.
  2. 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=True is not supported in that setup: there is no separate reference to update.
  3. 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=True does this, with precompute_ref_batch_size allowed to be larger than the training batch because the pass holds no activations for backward. It is not supported with streaming IterableDataset inputs, and it is incompatible with sync_ref_model, which moves the reference during training.
  4. 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 collapse

With 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/chosen steadily 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_norm exist 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" with label_smoothing set 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

  1. Start from your SFT checkpoint and a preference set in prompt, chosen and rejected format with an explicit prompt.
  2. Check chosen vs rejected length statistics and drop near-duplicate pairs before you spend GPU time.
  3. Run one step and confirm the first logged loss is about 0.693 and the reward accuracy is about 0.5.
  4. 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.
  5. Set model_init_kwargs dtype explicitly, keep log-prob sums in fp32 and chunk or fuse the loss before raising the context length.
  6. Track logps/chosen alongside margins, and stop when held-out preference accuracy or your capability suite stops improving.
Key takeaway: A DPO step scores each pair four times: policy and reference, on chosen and on rejected. It then backpropagates a logistic loss on the difference of log-ratios. On the GPU that means a second model's worth of inference, fp32-sensitive sequence sums and a vocabulary-sized logits tensor that often outgrows the layer activations. Make the reference cheap with adapter-disable or precomputed log-probabilities, chunk or fuse the log-softmax, shard full-parameter runs as you would SFT, and watch the chosen log-probability as well as the margin so the run doesn't quietly optimize itself into worse answers.