DPO and direct preference optimisation
Be able to derive DPO from the RLHF objective and fine-tune a model with preference pairs.
Prerequisites
- FRLHF and preference learningrequired
Intuition
RLHF needs three models (the policy, the reward model, the reference) and an RL loop that is hard to stabilise.
DPO's insight: the RL problem RLHF solves has a closed-form optimal solution. The optimal policy is the reference policy reweighted by the reward. Turn the relation around and the reward can be expressed in terms of the policy — and then the reward model disappears from the equation.
What is left is an ordinary classification loss on preference pairs: increase the probability of the chosen answer and decrease it for the rejected one, relative to the reference model.
No RL, no reward model, no sampling during training. Which is why DPO is standard in open models.
Derivation
The RLHF objective:
The solution is known (the Gibbs distribution):
Solve for the reward:
Substitute that into the Bradley–Terry model for preferences. The normalisation is the same for and and cancels:
Maximum likelihood on that gives the DPO loss. The gradient automatically upweights the examples where the model ranks them wrongly — it learns most from its mistakes.
Beta: a small β allows a larger deviation from the reference (more effect, a greater risk of degeneration); a large β keeps the policy close. 0.1 is a common start.
Variants: IPO (counteracts overfitting to the preferences), KTO (only needs good/bad, not pairs), ORPO (merges SFT and preference training into one step).
Code
import torch, torch.nn.functional as F
def sequence_logp(model, ids, prompt_len):
"""The summed log probability of the answer part (not the prompt)."""
out = model(ids).logits[:, :-1].log_softmax(-1)
target = ids[:, 1:]
logp = out.gather(-1, target.unsqueeze(-1)).squeeze(-1)
mask = torch.arange(target.size(1), device=ids.device)[None] >= (prompt_len - 1)
return (logp * mask).sum(-1)
def dpo_step(policy, ref, batch, beta=0.1, opt=None):
with torch.no_grad():
ref_w = sequence_logp(ref, batch["ids_w"], batch["plen"])
ref_l = sequence_logp(ref, batch["ids_l"], batch["plen"])
pol_w = sequence_logp(policy, batch["ids_w"], batch["plen"])
pol_l = sequence_logp(policy, batch["ids_l"], batch["plen"])
logits = beta * ((pol_w - ref_w) - (pol_l - ref_l))
loss = -F.logsigmoid(logits).mean()
loss.backward(); opt.step(); opt.zero_grad()
return {"loss": loss.item(),
"accuracy": (logits > 0).float().mean().item(), # how often the ranking is right
"margin": logits.mean().item(),
"kl_proxy_w": (pol_w - ref_w).mean().item()} # running away? overtrained
Diagnostics during training: accuracy should rise towards 0.7–0.9; kl_proxy should grow slowly. If the KL runs away while the accuracy stands still, the policy is on its way to degenerating — lower the lr or raise β.
Mastery means
- Derives DPO's idea from the RLHF objective
- Fine-tunes with preference pairs
- Chooses beta and detects an overtrained policy
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — Direct Preference Optimization — arXiv (open access; licence per article)
- arXiv — Training language models to follow instructions with human feedback — arXiv (open access; licence per article)