ml-explore / ml-explore/mlx-examples
Reinforcement Learning from Human Feedback (RLHF) examples: Direct Preference Optimization (DPO)
Nobody has claimed this yet.
- Dominant language
- Python
- Stars
- 9k
- Forks
- 1.2k
- PR merge metrics
- No merged PRs in 30d
Description
Introduce one Reinforcement Learning from Human Feedback (RLHF) example, such as Direct Preference Optimization (DPO) method.
Paper
Direct Preference Optimization: Your Language Model is Secretly a Reward Model
Notes
Direct Preference Optimization (DPO): A Simplified Explanation by João Lages

Implementation examples
- huggingface/trl: TRL - Transformer Reinforcement Learning
- eric-mitchell/direct-preference-optimization: Direct Preference Optimization
Possible MLX implementation
Policy and reference log probabilities:
def get_batched_logps(model, inputs, targets):
logits, _ = model(inputs)
logits = logits.astype(mx.float32)
loss_mask = targets != 0
per_token_logps = mx.take_along_axis(nn.log_softmax(logits), targets[..., None], axis=2).squeeze(2)
return tuple((per_token_logps * loss_mask).sum(-1).split(2))
Loss:
def dpo_loss(model, beta, label_smoothing, reference_chosen_logps, reference_rejected_logps, inputs, targets):
chosen_logps, rejected_logps = get_batched_logps(model, inputs, targets)
pi_logratios = chosen_logps - rejected_logps
reference_logratios = reference_chosen_logps - reference_rejected_logps
logits = pi_logratios - reference_logratios
losses = -nn.log_sigmoid(beta * logits) * (1.0 - label_smoothing) - nn.log_sigmoid(-beta * logits) * label_smoothing
chosen_rewards = beta * (chosen_logps - reference_chosen_logps)
rejected_rewards = beta * (rejected_logps - reference_rejected_logps)
reward_accuracies = (chosen_rewards > rejected_rewards).astype(mx.float32)
reward_margins = chosen_rewards - rejected_rewards
ntoks = (inputs != 0).sum()
return (
losses.mean(),
chosen_rewards.mean(),
rejected_rewards.mean(),
reward_accuracies.mean(),
reward_margins.mean(),
ntoks,
)
Beta: The temperature parameter for the DPO loss is typically set in the range of 0.1 to 0.5. The reference model is ignored when beta equals 0.
Label smoothing: This parameter represents the conservativeness for DPO loss, assuming that preferences are noisy and can be flipped with a probability of label_smoothing.
Note
label_smoothing > 0defines the Conservative DPO loss.
Contributor guide
First steps
- Read the whole issue, then the project's contributing guide.
- Comment on the issue to say you are picking it up — it saves two people doing the same work.
- Fork the repository and make your change on a branch.
- Open a pull request that references the issue number.
Research direction
Read the linked DPO paper and the referenced TRL and direct-preference-optimization implementations first. Use the provided get_batched_logps and dpo_loss sketches as the starting point for an MLX example, with beta and label_smoothing behavior covered; done means the requested DPO RLHF example is implemented and usable.
Written by the indexing model from the issue text.
Assessment
- Tech stack
- python
- Domain
- machine-learning
- Issue type
- Feature
- Difficulty
- 5/5
- Estimated time
- Over a week
- Activity status
- Stale
- Clarity
- Mostly clear
- Newbie friendliness
- 32/100