Training · After pretraining
- policy
- GPT-2 small + LoRA rank 4
- reference
- GPT-2 small, frozen
- β
- 0.1 25 steps
- rewards
- toy RLHF steps
All steps
Preferences, not answers
After instruction tuning, models are tuned on preferences: for one prompt, two answers and which one people preferred. Judging is easier than writing. In these three pairs GPT-2 small leans the other way: after “The cat sat on the” it finds “floor.” 21 times as likely as “mat.”.
prompt · chosen ≻ rejectedRLHF: a reward model
Classic RLHF first trains a reward model on such pairs: a network that reads a prompt and an answer and outputs one score, trained so that the preferred answer scores higher. Here it scores six possible next tokens after “The cat sat on the”; these scores are made up for the example.
P(chosen ≻ rejected) = σ(rchosen − rrejected)RLHF: sample, score, shift
Then reinforcement learning: the model samples an answer, the reward model scores it, and probability moves toward what scores well. A penalty on drifting from the start (the outlined bars) holds it back, so it settles at πref · exp(r / β) rather than on the single best token. Change β to loosen or tighten that pull.
max E[r] − β · KL(π ‖ πref)DPO: no reward model
DPO drops the reward model and the sampling. Its loss looks at each pair directly: how much more likely the model makes the chosen answer than the reference does, against the same for the rejected one, and it pushes that gap open. β again sets how far it may go.
−log σ(β [Δ log π(chosen) − Δ log π(rejected)])A real DPO run
A real run on GPT-2 small: rank-4 adapters, the frozen model as reference, β = 0.1, 25 steps on the three pairs. Each answer sits at its log-probability; the gap opens until the loss falls from 0.693 to 0.005. Most of the gap comes from the rejected answers falling.
loss 0.693 → 0.005Where the probability went
Where the probability went, for the next token after “The cat sat on the”: GPT-2 spread it over floor, bed, couch and more; after 25 steps nearly all of it is on “mat”. With three pairs the model simply memorized them. Real runs use many thousands of pairs, and β and the reference keep them from drifting this far.
rejected pushed down more than chosen pulled up
Code
def logp(model, prompt, answer): # Σ log p of the answer tokens
logits = model(prompt + answer).logits[len(prompt) - 1:-1]
return logits.log_softmax(-1).gather(-1, answer[:, None]).sum()
r_w = beta * (logp(policy, x, y_w) - logp(ref, x, y_w)) # implicit rewards
r_l = beta * (logp(policy, x, y_l) - logp(ref, x, y_l))
loss = -F.logsigmoid(r_w - r_l)
loss.backward(); opt.step() # ref is frozen; only the policy moves