DPO och direkt preferensoptimering
Kunna härleda DPO ur RLHF-målet och finjustera en modell med preferenspar.
Förkunskaper
Intuition
RLHF kräver tre modeller (policy, belöningsmodell, referens) och en RL-loop som är svår att stabilisera.
DPO:s insikt: det RL-problem RLHF löser har en sluten optimal lösning. Den optimala policyn är referenspolicyn omviktad med belöningen. Vänder man på sambandet kan belöningen uttryckas i policyn — och då försvinner belöningsmodellen ur ekvationen.
Kvar blir en vanlig klassificeringsförlust på preferenspar: öka sannolikheten för det valda svaret och minska för det förkastade, relativt referensmodellen.
Ingen RL, ingen belöningsmodell, ingen sampling under träning. Därför är DPO standard i öppna modeller.
Härledning
RLHF-målet:
Lösningen är känd (Gibbs-fördelningen):
Lös ut belöningen:
Sätt in i Bradley–Terry-modellen för preferenser. Normaliseringen är samma för och och kancellerar:
Maximum likelihood på det ger DPO-förlusten. Gradienten viktar automatiskt upp exempel där modellen rangordnar fel — den lär sig mest av sina misstag.
Beta: litet β tillåter större avvikelse från referensen (mer effekt, större risk för degenerering); stort β håller policyn nära. 0,1 är vanlig start.
Varianter: IPO (motverkar överanpassning till preferenser), KTO (behöver bara bra/dåligt, inte par), ORPO (slår ihop SFT och preferensträning i ett steg).
Kod
import torch, torch.nn.functional as F
def sekvens_logp(modell, ids, prompt_len):
"""Summerad log-sannolikhet för svarsdelen (inte prompten)."""
ut = modell(ids).logits[:, :-1].log_softmax(-1)
mal = ids[:, 1:]
logp = ut.gather(-1, mal.unsqueeze(-1)).squeeze(-1)
mask = torch.arange(mal.size(1), device=ids.device)[None] >= (prompt_len - 1)
return (logp * mask).sum(-1)
def dpo_steg(policy, ref, batch, beta=0.1, opt=None):
with torch.no_grad():
ref_w = sekvens_logp(ref, batch["ids_w"], batch["plen"])
ref_l = sekvens_logp(ref, batch["ids_l"], batch["plen"])
pol_w = sekvens_logp(policy, batch["ids_w"], batch["plen"])
pol_l = sekvens_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(),
"traffsakerhet": (logits > 0).float().mean().item(), # hur ofta rangordnas rätt
"marginal": logits.mean().item(),
"kl_proxy_w": (pol_w - ref_w).mean().item()} # skenar den? övertränad
Diagnostik under träning: traffsakerhet ska stiga mot 0,7–0,9; kl_proxy ska växa långsamt. Skenar KL medan träffsäkerheten står still är policyn på väg att degenerera — sänk lr eller höj β.
Behärskning innebär
- Härleder DPO:s idé ur RLHF-målet
- Finjusterar med preferenspar
- Väljer beta och upptäcker övertränad policy
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- arXiv — Direct Preference Optimization — arXiv (öppen åtkomst; licens per artikel)
- arXiv — Training language models to follow instructions with human feedback — arXiv (öppen åtkomst; licens per artikel)