当前位置:首页 > 技术 > 正文内容

DPO与PPO算法原理及PyTorch实现对比

访客 技术 2026年7月22日 3

概述

在大型语言模型(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。
  • 两者均可提升模型对齐能力,具体选择应基于实际资源和需求。

相关文章

Linux crontab 详解

1) crontab 是什么cron 是 Linux 的定时任务守护进程;crontab 是用来编辑/查看“按时间周期执行命令”的表(cron table)。常见两类:用户 crontab:每个用户一份(crontab -e 编辑)系统级 crontab / cron.d:可指定执行用户(/etc/crontab、/etc/cron.d/*)2) crontab 时间...

富文本里可以允许的 HTML 属性

一、所有标签默认允许的安全属性(极少)class        (可选)id           (通常建议禁用)title️ 注意:id 容易被滥用做锚点注入,很多系统直接禁用class 允许的话最好只允许固定前缀(如 editor-*)二、a 标签允许属性<a href="" t...

Mac 安装 Node.js 指南

方法一:通过官网安装包(最简单,适合初学者)如果你只是想快速安装并开始使用,这是最直接的方法。访问 Node.js 官网。页面会显示两个版本:LTS (Recommended For Most Users):长期支持版,最稳定。建议选这个。Current:最新特性版,包含最新功能但可能不够稳定。下载 .pkg 安装包并运行。按照安装向导点击“下一步”即可完成。方法二:使用 Homebrew 安装(...

Dom\HTML_NO_DEFAULT_NS 的副作用:自动加闭合标签

在使用Dom\HTMLDocument时,Dom\HTML_NO_DEFAULT_NS 将禁止在解析过程中设置元素的命名空间, 此设置是为了与DOMDocument向后兼容而存在的。当使用它时,已知的一个副作用就是:自动加闭合标签例如 </img> 为什么会这样?当你使用:Dom\HTML_NO_DEFAULT_NS文档会变成 无命名空间模式,此时内部更接近 XML...

Laravel 事件和监听器创建

在 Laravel 中,使用 Artisan 命令创建 Events(事件) 和 Listeners(监听器) 是非常高效的。你可以通过以下几种方式来实现:1. 手动创建单个 Event如果你只想创建一个事件类,可以使用 make:event 命令:Bashphp artisan make:event UserRegistered执行后,文件将生成在 app/Even...

自定义域名解析神器 dnsmasq

什么是 dnsmasq?dnsmasq 是一个轻量级、功能强大的网络服务工具,专为小型和中等规模网络设计。它是一个综合的网络基础设施解决方案[1]。dnsmasq 能做什么?功能说明应用场景DNS 转发与缓存将 DNS 查询转发到上游服务器(ISP、Google DNS 等),并在本地缓存结果加快 DNS 查询速度,减少外部 DNS 流量本地 DNS解析本地网络设备的主机名,无需编辑&n...

发表评论

访客

◎欢迎参与讨论,请在这里发表您的看法和观点。