欢迎光临
我们一直在努力

DPO直接偏好优化算法的理论研究和实现

目录

1.DPO基础建模

2.DPO奖励函数

3. DPO的损失函数

4.Python代码实现


       基于近端策略优化(PPO)的人类反馈强化学习(RLHF)凭借其在ChatGPT等模型上的表现,成为了对齐训练的主流范式。然而,RLHF复杂的训练流程、对强化学习(RL)专业知识的高度依赖,以及训练过程中难以避免的不稳定性,都成为了其规模化落地的瓶颈。而直接偏好优化(Direct Preference Optimization, DPO)算法应运而生。由斯坦福大学等机构的研究团队提出,以其“直接优化策略,无需显式奖励模型”的核心思想,彻底简化了对齐训练的流程,将复杂的强化学习问题转化为了直观的监督学习问题。

1.DPO基础建模

DPO的数学推导始于对偏好数据的建模。我们首先定义一个偏好数据集:

其中:

       在RLHF中,我们会训练一个显式的奖励模型rϕ(x,y),其目标是拟合人类对回答的偏好。这个拟合过程可以被建模为一个二分类问题:我们希望奖励模型能够正确判断,在给定Prompt x的情况下,回答yw比yl更优。其损失函数为负对数似然损失:

2.DPO奖励函数

       DPO没有像RLHF那样去显式地训练一个奖励模型rϕ,而是提出了一个隐式奖励函数r(x,y)。这个函数并非一个独立的神经网络,而是通过策略模型πθ和参考模型πref的概率对数比来定义的:

其中:

它将奖励的概念内生于策略模型本身。一个回答y的“奖励”,不再是由一个外部模型评判,而是由它在策略模型和参考模型下的相对概率来决定。如果策略模型生成y的概率远高于参考模型,那么y就获得了高奖励;反之则获得低奖励。

在DPO的原论文中,这个隐式奖励函数有一个更完整的推导形式,包含了一个配分项Z(x):

由于DPO的优化目标是基于Bradley-Terry建模理念,该理念仅依赖于两项奖励值之间的差异,而非奖励的绝对值。因此,在计算r(x,yw)−r(x,yl)时,式中的未知项βlogZ(x)会被抵消。

3. DPO的损失函数

       将DPO定义的隐式奖励函数代入到RLHF中奖励模型的损失函数中,就可以直接得到DPO的损失函数LDPO:

其优化目标J(θ)就是最大化这个期望:

其中:

通过这种方式,DPO在训练过程中,持续地引导策略模型πθ,使其在给定输入x的情况下,生成高质量回答yw的可能性越来越大,同时生成低质量回答yl的可能性越来越小,最终实现与人类偏好的对齐。

4.Python代码实现

import torch
import torch.nn as nn
import torch.optim as optim
import torch.nn.functional as F
from torch.utils.data import Dataset, DataLoader
import numpy as np

# 设置随机种子以保证结果可复现
torch.manual_seed(42)
np.random.seed(42)

# 设备配置:优先使用GPU,如果没有则使用CPU
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")

# ======================
# 1. 定义简化的语言模型(策略模型/参考模型)
# ======================
class SimpleLM(nn.Module):
"""
简化的语言模型,用于演示DPO训练过程
实际应用中会替换为真实的大语言模型(如LLaMA、GPT等)
"""
def __init__(self, vocab_size=1000, embed_dim=128, hidden_dim=256):
super(SimpleLM, self).__init__()
self.embedding = nn.Embedding(vocab_size, embed_dim)
self.lstm = nn.LSTM(embed_dim, hidden_dim, batch_first=True)
self.fc = nn.Linear(hidden_dim, vocab_size)

def forward(self, x):
"""
前向传播,返回每个位置的logits
Args:
x: 输入序列,shape [batch_size, seq_len]
Returns:
logits: 输出logits,shape [batch_size, seq_len, vocab_size]
"""
embed = self.embedding(x) # [batch_size, seq_len, embed_dim]
lstm_out, _ = self.lstm(embed) # [batch_size, seq_len, hidden_dim]
logits = self.fc(lstm_out) # [batch_size, seq_len, vocab_size]
return logits

def get_log_prob(self, x, y):
"""
计算模型生成回答y的对数概率
Args:
x: 输入prompt,shape [batch_size, seq_len_x]
y: 回答序列,shape [batch_size, seq_len_y]
Returns:
log_prob: 对数概率,shape [batch_size]
"""
# 获取模型输出logits
logits = self.forward(torch.cat([x, y], dim=1)) # 拼接prompt和回答

# 只取回答部分的logits(去掉prompt部分)
logits = logits[:, x.shape[1]:-1, :] # 去掉最后一个token,因为要预测下一个

# 计算对数概率
log_probs = F.log_softmax(logits, dim=-1)

# 取出每个token对应的对数概率
target_log_probs = log_probs.gather(
dim=2,
index=y[:, 1:].unsqueeze(-1) # 回答从第二个token开始(去掉<BOS>)
).squeeze(-1)

# 对序列长度维度求和,得到整个回答的对数概率
return target_log_probs.sum(dim=1)

# ======================
# 2. 实现DPO损失函数
# ======================
class DPOLoss(nn.Module):
"""
DPO损失函数实现
核心公式:L_DPO = -logσ(β * (logπθ(y_w|x)/π_ref(y_w|x) – logπθ(y_l|x)/π_ref(y_l|x)))
"""
def __init__(self, beta=0.1):
super(DPOLoss, self).__init__()
self.beta = beta # DPO的温度系数

def forward(self, policy_model, ref_model, x, y_w, y_l):
"""
计算DPO损失
Args:
policy_model: 待优化的策略模型
ref_model: 参考模型(通常是SFT模型)
x: 输入prompt,shape [batch_size, seq_len_x]
y_w: 优质回答,shape [batch_size, seq_len_yw]
y_l: 劣质回答,shape [batch_size, seq_len_yl]
Returns:
loss: DPO损失值
rewards: 包含各项奖励值的字典
"""
# 1. 计算策略模型对优质/劣质回答的对数概率
policy_logp_w = policy_model.get_log_prob(x, y_w) # [batch_size]
policy_logp_l = policy_model.get_log_prob(x, y_l) # [batch_size]

# 2. 计算参考模型对优质/劣质回答的对数概率
with torch.no_grad(): # 参考模型不参与梯度更新
ref_logp_w = ref_model.get_log_prob(x, y_w) # [batch_size]
ref_logp_l = ref_model.get_log_prob(x, y_l) # [batch_size]

# 3. 计算优势(advantage): r_w – r_l
# r = β * log(πθ(y|x)/π_ref(y|x)) = β * (logπθ(y|x) – logπ_ref(y|x))
r_w = self.beta * (policy_logp_w – ref_logp_w)
r_l = self.beta * (policy_logp_l – ref_logp_l)
advantage = r_w – r_l # [batch_size]

# 4. 计算DPO损失
loss = -F.logsigmoid(advantage).mean()

# 记录奖励值,用于监控训练过程
rewards = {
"reward_win": r_w.mean().item(),
"reward_lose": r_l.mean().item(),
"advantage": advantage.mean().item()
}

return loss, rewards

# ======================
# 3. 构建偏好数据集
# ======================
class PreferenceDataset(Dataset):
"""
偏好数据集,包含prompt、优质回答、劣质回答
"""
def __init__(self, num_samples=1000, seq_len_x=10, seq_len_y=20, vocab_size=1000):
self.num_samples = num_samples
self.seq_len_x = seq_len_x
self.seq_len_y = seq_len_y
self.vocab_size = vocab_size

# 生成模拟数据
self.data = self._generate_data()

def _generate_data(self):
"""生成模拟的偏好数据"""
data = []
for _ in range(self.num_samples):
# 生成prompt (x)
x = torch.randint(1, self.vocab_size, (self.seq_len_x,))

# 生成优质回答 (y_w) – 模式更固定
y_w = torch.cat([
torch.tensor([0]), # <BOS> token
torch.randint(1, 500, (self.seq_len_y-1,)) # 前500个token作为优质区
])

# 生成劣质回答 (y_l) – 模式更随机
y_l = torch.cat([
torch.tensor([0]), # <BOS> token
torch.randint(500, self.vocab_size, (self.seq_len_y-1,)) # 后500个token作为劣质区
])

data.append((x, y_w, y_l))
return data

def __len__(self):
return self.num_samples

def __getitem__(self, idx):
return self.data[idx]

# 数据加载器collate函数
def collate_fn(batch):
"""将batch数据整理成tensor"""
xs, y_ws, y_ls = zip(*batch)
x = torch.stack(xs).to(device)
y_w = torch.stack(y_ws).to(device)
y_l = torch.stack(y_ls).to(device)
return x, y_w, y_l

# ======================
# 4. 训练流程
# ======================
def train_dpo():
# 超参数设置
vocab_size = 1000
embed_dim = 128
hidden_dim = 256
batch_size = 32
epochs = 20
lr = 1e-4
beta = 0.1 # DPO温度系数

# 1. 创建模型
# 策略模型(待优化)
policy_model = SimpleLM(vocab_size, embed_dim, hidden_dim).to(device)
# 参考模型(SFT模型,固定参数)
ref_model = SimpleLM(vocab_size, embed_dim, hidden_dim).to(device)

# 固定参考模型参数(不参与训练)
for param in ref_model.parameters():
param.requires_grad = False

# 2. 创建数据集和数据加载器
dataset = PreferenceDataset(num_samples=1000, vocab_size=vocab_size)
dataloader = DataLoader(
dataset,
batch_size=batch_size,
shuffle=True,
collate_fn=collate_fn
)

# 3. 初始化损失函数和优化器
dpo_loss_fn = DPOLoss(beta=beta)
optimizer = optim.AdamW(policy_model.parameters(), lr=lr)

# 4. 训练循环
policy_model.train()
for epoch in range(epochs):
total_loss = 0.0
avg_reward_win = 0.0
avg_reward_lose = 0.0
avg_advantage = 0.0

for batch_idx, (x, y_w, y_l) in enumerate(dataloader):
optimizer.zero_grad()

# 计算DPO损失
loss, rewards = dpo_loss_fn(policy_model, ref_model, x, y_w, y_l)

# 反向传播和优化
loss.backward()
optimizer.step()

# 累计统计
total_loss += loss.item()
avg_reward_win += rewards["reward_win"]
avg_reward_lose += rewards["reward_lose"]
avg_advantage += rewards["advantage"]

# 计算平均指标
avg_loss = total_loss / len(dataloader)
avg_reward_win /= len(dataloader)
avg_reward_lose /= len(dataloader)
avg_advantage /= len(dataloader)

# 打印训练信息
print(f"Epoch [{epoch+1}/{epochs}]")
print(f" Loss: {avg_loss:.4f}")
print(f" Reward (win): {avg_reward_win:.4f}")
print(f" Reward (lose): {avg_reward_lose:.4f}")
print(f" Advantage: {avg_advantage:.4f}")
print("-" * 50)

return policy_model, ref_model

# ======================
# 5. 验证训练效果
# ======================
def validate_model(policy_model, ref_model):
"""验证DPO训练后的模型效果"""
# 创建测试数据
vocab_size = 1000
x = torch.randint(1, vocab_size, (1, 10)).to(device) # 测试prompt
y_w = torch.cat([torch.tensor([[0]]), torch.randint(1, 500, (1, 19))], dim=1).to(device) # 优质回答
y_l = torch.cat([torch.tensor([[0]]), torch.randint(500, vocab_size, (1, 19))], dim=1).to(device) # 劣质回答

# 计算策略模型的概率
policy_logp_w = policy_model.get_log_prob(x, y_w)
policy_logp_l = policy_model.get_log_prob(x, y_l)

# 计算参考模型的概率
ref_logp_w = ref_model.get_log_prob(x, y_w)
ref_logp_l = ref_model.get_log_prob(x, y_l)

# 打印结果
print("\\n=== 训练效果验证 ===")
print(f"策略模型对优质回答的对数概率: {policy_logp_w.item():.4f}")
print(f"策略模型对劣质回答的对数概率: {policy_logp_l.item():.4f}")
print(f"参考模型对优质回答的对数概率: {ref_logp_w.item():.4f}")
print(f"参考模型对劣质回答的对数概率: {ref_logp_l.item():.4f}")

# 计算概率比
policy_ratio = torch.exp(policy_logp_w – policy_logp_l).item()
ref_ratio = torch.exp(ref_logp_w – ref_logp_l).item()

print(f"\\n策略模型: 优质回答概率 / 劣质回答概率 = {policy_ratio:.2f}")
print(f"参考模型: 优质回答概率 / 劣质回答概率 = {ref_ratio:.2f}")
print(f"\\n训练后,策略模型更偏好优质回答的程度提升了 {policy_ratio/ref_ratio:.2f} 倍")

# ======================
# 主函数
# ======================
if __name__ == "__main__":
# 训练DPO模型
print("开始DPO训练…")
policy_model, ref_model = train_dpo()

# 验证训练效果
validate_model(policy_model, ref_model)

赞(0)
未经允许不得转载:171主机测评 » DPO直接偏好优化算法的理论研究和实现
分享到: 更多 (0)

评论 抢沙发

  • 昵称 (必填)
  • 邮箱 (必填)
  • 网址