DPO与PPO算法原理及PyTorch实现对比
概述
在大型语言模型(LLM)训练中,强化学习(RL)常被用于对齐模型输出与人类偏好。PPO(近端策略优化)是RLHF(基于人类反馈的强化学习)中的经典算法,而DPO(直接偏好优化)是一种无需奖励模型的新兴方法。本文将从原理和代码层面对比两者,帮助理解其核心思想与工程实践。
1. PPO算法原理
PPO是一种策略梯度算法,通过限制新旧策略的偏差来稳定训练。在RLHF流程中,它利用奖励模型评估模型输出,从而优化策略。
核心损失函数
PPO的损失函数为:
L_CLIP(θ) = E[min( r(θ)A, clip(r(θ), 1-ε, 1+ε)A )]
其中,r(θ)是策略概率比,A是优势函数,ε是裁剪范围。
PyTorch实现(简化版)
import torch
import torch.nn as nn
import torch.optim as optim
class PolicyNetwork(nn.Module):
def __init__(self, state_dim, action_dim):
super().__init__()
self.net = nn.Sequential(
nn.Linear(state_dim, 64),
nn.ReLU(),
nn.Linear(64, action_dim)
)
def forward(self, state):
return self.net(state)
def ppo_loss(new_policy, old_policy, states, actions, advantages, clip_eps=0.2):
new_logits = new_policy(states)
new_dist = torch.distributions.Categorical(logits=new_logits)
new_log_probs = new_dist.log_prob(actions)
with torch.no_grad():
old_logits = old_policy(states)
old_dist = torch.distributions.Categorical(logits=old_logits)
old_log_probs = old_dist.log_prob(actions)
ratio = torch.exp(new_log_probs - old_log_probs)
clipped_advantages = torch.clamp(ratio, 1 - clip_eps, 1 + clip_eps) * advantages
loss = -torch.mean(torch.min(ratio * advantages, clipped_advantages))
return loss
# 训练循环示例
policy = PolicyNetwork(4, 2)
old_policy = PolicyNetwork(4, 2)
optimizer = optim.Adam(policy.parameters(), lr=0.001)
for epoch in range(100):
# 假设已收集数据:states, actions, advantages
loss = ppo_loss(policy, old_policy, states, actions, advantages)
optimizer.zero_grad()
loss.backward()
optimizer.step()
old_policy.load_state_dict(policy.state_dict())
2. DPO算法原理
DPO直接利用人类偏好对比数据(优选和次选输出)优化模型,无需训练奖励模型。它通过最大化优选输出相对于次选输出的对数概率差异来学习。
核心损失函数
DPO的损失函数为:
L_DPO = -E[ log( sigmoid( β * (logπ(y_win|x) - logπ(y_lose|x)) ) ) ]
其中,y_win是优选输出,y_lose是次选输出,β是温度参数。
PyTorch实现(简化版)
class SimpleLM(nn.Module):
def __init__(self, vocab_size, embed_dim):
super().__init__()
self.emb = nn.Embedding(vocab_size, embed_dim)
self.out = nn.Linear(embed_dim, vocab_size)
def forward(self, input_ids):
x = self.emb(input_ids).mean(dim=1)
return self.out(x)
def dpo_loss(model, prompt, win_response, lose_response, beta=0.1):
win_logits = model(torch.cat([prompt, win_response], dim=1))
lose_logits = model(torch.cat([prompt, lose_response], dim=1))
win_log_probs = torch.log_softmax(win_logits, dim=-1).gather(1, win_response).sum(dim=1)
lose_log_probs = torch.log_softmax(lose_logits, dim=-1).gather(1, lose_response).sum(dim=1)
diff = beta * (win_log_probs - lose_log_probs)
loss = -torch.log(torch.sigmoid(diff)).mean()
return loss
# 训练循环示例
model = SimpleLM(5000, 256)
optimizer = optim.Adam(model.parameters(), lr=1e-4)
for epoch in range(50):
# 假设已准备数据:prompt, win_response, lose_response
loss = dpo_loss(model, prompt, win_response, lose_response)
optimizer.zero_grad()
loss.backward()
optimizer.step()
3. 对比分析
| 维度 | PPO | DPO |
|---|---|---|
| 训练流程 | 需要奖励模型,流程复杂 | 直接使用偏好数据,流程简单 |
| 实现难度 | 较高,需要采样和优势计算 | 较低,损失函数直接 |
| 稳定性 | 依赖奖励模型质量 | 超参数少,更稳定 |
| 成熟度 | 已广泛使用,有成熟的库支持 | 较新,但已被多家采用 |
4. 选择建议
- 若已有高质量的奖励模型,或需要更传统的RLHF流程,选择PPO。
- 若只有成对人类偏好数据,且希望实现更简洁,选择DPO。
- 两者均可提升模型对齐能力,具体选择应基于实际资源和需求。