PPO GRPO GSPO DAPO的Loss计算与代码实现

📅 2026/7/25 0:35:00 👁️ 阅读次数 📝 编程学习
PPO GRPO GSPO DAPO的Loss计算与代码实现

PPO、GRPO、GSPO、DAPO 的 Loss 计算与代码实现

在强化学习(Reinforcement Learning, RL)领域,策略优化算法一直是研究的核心。从经典的 PPO(Proximal Policy Optimization)到近年来出现的 GRPO(Group Relative Policy Optimization)、GSPO(Generalized Surrogate Policy Optimization)以及 DAPO(Dual-Agent Policy Optimization),这些算法通过不同的 Loss 设计,解决了策略更新中的稳定性、样本效率以及多智能体协作等问题。本文将深入剖析这四种算法的 Loss 计算原理,并提供可运行的代码片段,帮助读者从底层理解其工作机制。## PPO:基于信任区域的策略优化PPO(Proximal Policy Optimization)由 OpenAI 在 2017 年提出,其核心思想是通过裁剪(Clipping)机制,限制策略更新的幅度,避免因单步更新过大导致性能崩溃。PPO 的 Loss 通常包含三部分:策略损失(Policy Loss)、价值损失(Value Loss)和熵正则项(Entropy Bonus)。### Loss 计算原理PPO 的策略损失基于重要性采样(Importance Sampling)和裁剪:[L^{CLIP}(\theta) = \mathbb{E}t \left[ \min\left( r_t(\theta) \hat{A}t, \text{clip}(r_t(\theta), 1-\epsilon, 1+\epsilon) \hat{A}t \right) \right]]其中,( r_t(\theta) = \frac{\pi\theta(a_t|s_t)}{\pi{\theta{old}}(a_t|s_t)} ) 是重要性权重,( \hat{A}_t ) 是优势函数估计,( \epsilon ) 是裁剪阈值(通常为 0.2)。价值损失通常使用均方误差(MSE)计算:( L^{VF}(\theta) = \mathbb{E}t[(V\theta(s_t) - R_t)^2] ),其中 ( R_t ) 是折扣回报。### 代码实现以下是一个简化且可运行的 PPO Loss 计算代码片段:pythonimport torchimport torch.nn as nndef ppo_loss(old_log_probs, new_log_probs, advantages, values, returns, epsilon=0.2, entropy_coef=0.01, value_coef=0.5): """ 计算 PPO 的 Loss :param old_log_probs: 旧策略的 log 概率 (tensor) :param new_log_probs: 新策略的 log 概率 (tensor) :param advantages: 优势函数 (tensor) :param values: 价值函数预测值 (tensor) :param returns: 折扣回报 (tensor) :param epsilon: 裁剪阈值 :param entropy_coef: 熵正则系数 :param value_coef: 价值损失系数 :return: 总损失 (tensor) """ # 1. 计算重要性权重 ratio ratio = torch.exp(new_log_probs - old_log_probs) # r_t(theta) # 2. 无裁剪的 surrogate loss surr1 = ratio * advantages # 3. 裁剪后的 surrogate loss surr2 = torch.clamp(ratio, 1.0 - epsilon, 1.0 + epsilon) * advantages # 4. 策略损失:取最小值以限制更新 policy_loss = -torch.min(surr1, surr2).mean() # 5. 价值损失(MSE) value_loss = nn.MSELoss()(values, returns) # 6. 熵正则项(鼓励探索) entropy = -(torch.exp(new_log_probs) * new_log_probs).mean() # 7. 总损失 total_loss = policy_loss + value_coef * value_loss - entropy_coef * entropy return total_loss, policy_loss, value_loss, entropy# 示例数据old_log_probs = torch.tensor([-0.5, -1.2, -0.8], requires_grad=False)new_log_probs = torch.tensor([-0.3, -1.0, -0.6], requires_grad=True)advantages = torch.tensor([1.0, -0.5, 0.8])values = torch.tensor([0.9, 0.3, 0.7], requires_grad=True)returns = torch.tensor([1.2, 0.1, 0.9])loss, p_loss, v_loss, ent = ppo_loss(old_log_probs, new_log_probs, advantages, values, returns)print(f"PPO Total Loss: {loss.item():.4f}, Policy Loss: {p_loss.item():.4f}, Value Loss: {v_loss.item():.4f}")## GRPO:群体相对策略优化GRPO(Group Relative Policy Optimization)是一种在多智能体强化学习(MARL)中提出的变体,其核心是将策略更新与群体内其他智能体的表现进行相对比较。GRPO 通过群体优势函数(Group Advantage)来调整每个智能体的 Loss,从而促进协作或竞争。### Loss 计算原理GRPO 的 Loss 定义如下:[L^{GRPO}(\theta_i) = \mathbb{E}_t \left[ \min\left( r_t(\theta_i) \hat{A}t^i, \text{clip}(r_t(\theta_i), 1-\epsilon, 1+\epsilon) \hat{A}t^i \right) \right] + \beta \cdot \text{KL}(\pi{\theta_i} | \pi{\text{group}})]其中,( \hat{A}t^i ) 是智能体 i 的群体优势函数,通常定义为 ( \hat{A}t^i = R_t^i - \frac{1}{N}\sum{j=1}^N R_t^j ),即个体回报与群体平均回报的差值。KL 散度项用于控制策略与群体策略的差异。### 代码实现以下是一个 GRPO Loss 的计算示例:pythonimport torchimport torch.nn as nnimport torch.nn.functional as Fdef grpo_loss(old_log_probs, new_log_probs, rewards, group_rewards, epsilon=0.2, beta=0.01): """ 计算 GRPO 的 Loss :param old_log_probs: 旧策略的 log 概率 (tensor, shape=[batch, n_agents]) :param new_log_probs: 新策略的 log 概率 (tensor, shape=[batch, n_agents]) :param rewards: 每个智能体的回报 (tensor, shape=[batch, n_agents]) :param group_rewards: 群体平均回报 (tensor, shape=[batch, 1]) :param epsilon: 裁剪阈值 :param beta: KL 散度系数 :return: 总损失 (tensor) """ # 1. 计算群体优势函数:个体回报减去群体平均 advantages = rewards - group_rewards # shape: [batch, n_agents] # 2. 计算重要性权重 ratio = torch.exp(new_log_probs - old_log_probs) # 3. 裁剪 surrogate loss surr1 = ratio * advantages surr2 = torch.clamp(ratio, 1.0 - epsilon, 1.0 + epsilon) * advantages policy_loss = -torch.min(surr1, surr2).mean() # 4. KL 散度正则项:衡量与群体策略的差异 # 假设群体策略 log prob 为 old_log_probs 的均值 group_log_probs = old_log_probs.mean(dim=1, keepdim=True).expand_as(old_log_probs) kl_div = F.kl_div(new_log_probs, group_log_probs, reduction='batchmean', log_target=True) # 5. 总损失 total_loss = policy_loss + beta * kl_div return total_loss, policy_loss, kl_div# 示例数据:2 个智能体,3 个时间步batch_size, n_agents = 3, 2old_log_probs = torch.tensor([[-0.5, -1.2], [-0.8, -0.3], [-1.0, -0.6]])new_log_probs = torch.tensor([[-0.3, -1.0], [-0.6, -0.1], [-0.8, -0.4]], requires_grad=True)rewards = torch.tensor([[1.0, 0.5], [0.8, 1.2], [0.3, 0.7]])group_rewards = rewards.mean(dim=1, keepdim=True) # 群体平均loss, p_loss, kl = grpo_loss(old_log_probs, new_log_probs, rewards, group_rewards)print(f"GRPO Total Loss: {loss.item():.4f}, Policy Loss: {p_loss.item():.4f}, KL Div: {kl.item():.4f}")## GSPO:广义替代策略优化GSPO(Generalized Surrogate Policy Optimization)是对 PPO 的推广,它引入了更灵活的替代目标函数,允许使用不同的距离度量(如 KL 散度、Fisher 信息矩阵)来约束策略更新。GSPO 的核心是将策略优化问题形式化为一个带约束的优化,并通过拉格朗日乘子法求解。### Loss 计算原理GSPO 的 Loss 形式为:[L^{GSPO}(\theta) = \mathbb{E}t \left[ r_t(\theta) \hat{A}t \right] - \lambda \cdot D(\pi\theta | \pi{\theta{old}})]其中,( D(\cdot | \cdot) ) 是一个距离函数(例如 KL 散度),( \lambda ) 是自适应调整的惩罚系数。与 PPO 的硬裁剪不同,GSPO 使用软约束。### 代码实现pythonimport torchimport torch.nn.functional as Fdef gspo_loss(old_log_probs, new_log_probs, advantages, lambda_coef=0.1, distance='kl'): """ 计算 GSPO 的 Loss :param old_log_probs: 旧策略 log 概率 (tensor) :param new_log_probs: 新策略 log 概率 (tensor) :param advantages: 优势函数 (tensor) :param lambda_coef: 惩罚系数 :param distance: 距离度量类型 ('kl' 或 'js') :return: 总损失 (tensor) """ # 1. 重要性采样目标 ratio = torch.exp(new_log_probs - old_log_probs) surrogate = (ratio * advantages).mean() # 2. 计算距离正则项 if distance == 'kl': # KL 散度:D_KL(π_new || π_old) kl_div = torch.mean(torch.exp(old_log_probs) * (old_log_probs - new_log_probs)) elif distance == 'js': # Jensen-Shannon 散度(对称版本) m_log_probs = 0.5 * (torch.exp(new_log_probs) + torch.exp(old_log_probs)).log() kl1 = F.kl_div(new_log_probs, m_log_probs, reduction='batchmean', log_target=True) kl2 = F.kl_div(old_log_probs, m_log_probs, reduction='batchmean', log_target=True) js_div = 0.5 * (kl1 + kl2) kl_div = js_div else: raise ValueError("Unsupported distance metric") # 3. 总损失:最大化 surrogate,最小化距离 total_loss = -surrogate + lambda_coef * kl_div return total_loss, surrogate, kl_div# 示例数据old_log_probs = torch.tensor([-0.5, -1.2, -0.8])new_log_probs = torch.tensor([-0.3, -1.0, -0.6], requires_grad=True)advantages = torch.tensor([1.0, -0.5, 0.8])loss, surr, kl = gspo_loss(old_log_probs, new_log_probs, advantages, lambda_coef=0.5, distance='kl')print(f"GSPO Total Loss: {loss.item():.4f}, Surrogate: {surr.item():.4f}, KL: {kl.item():.4f}")## DAPO:双智能体策略优化DAPO(Dual-Agent Policy Optimization)是一种针对双智能体或对抗性环境的算法,它通过引入一个辅助智能体(如对手或合作者)来调整主智能体的策略。DAPO 的 Loss 通常包含主策略损失和辅助策略损失的耦合项。### Loss 计算原理DAPO 的 Loss 定义为:[L^{DAPO}(\theta_m, \theta_a) = \mathbb{E}_t \left[ \min\left( r_t(\theta_m) \hat{A}_t^m, \text{clip}(r_t(\theta_m), 1-\epsilon, 1+\epsilon) \hat{A}_t^m \right) \right] + \alpha \cdot L^{aux}(\theta_a)]其中,( \theta_m ) 是主智能体策略参数,( \theta_a ) 是辅助智能体策略参数,( L^{aux} ) 可以是辅助智能体的 PPO 损失或探索奖励。### 代码实现pythonimport torchdef dapo_loss(main_old_log_probs, main_new_log_probs, aux_old_log_probs, aux_new_log_probs, main_advantages, aux_advantages, alpha=0.5, epsilon=0.2): """ 计算 DAPO 的 Loss :param main_old_log_probs: 主智能体旧策略 log 概率 (tensor) :param main_new_log_probs: 主智能体新策略 log 概率 (tensor) :param aux_old_log_probs: 辅助智能体旧策略 log 概率 (tensor) :param aux_new_log_probs: 辅助智能体新策略 log 概率 (tensor) :param main_advantages: 主智能体优势函数 (tensor) :param aux_advantages: 辅助智能体优势函数 (tensor) :param alpha: 辅助损失权重 :param epsilon: 裁剪阈值 :return: 总损失 (tensor) """ # 主智能体 PPO 损失 ratio_main = torch.exp(main_new_log_probs - main_old_log_probs) surr1 = ratio_main * main_advantages surr2 = torch.clamp(ratio_main, 1.0 - epsilon, 1.0 + epsilon) * main_advantages main_loss = -torch.min(surr1, surr2).mean() # 辅助智能体 PPO 损失(例如,对手策略) ratio_aux = torch.exp(aux_new_log_probs - aux_old_log_probs) surr1_aux = ratio_aux * aux_advantages surr2_aux = torch.clamp(ratio_aux, 1.0 - epsilon, 1.0 + epsilon) * aux_advantages aux_loss = -torch.min(surr1_aux, surr2_aux).mean() # 总损失 total_loss = main_loss + alpha * aux_loss return total_loss, main_loss, aux_loss# 示例数据main_old = torch.tensor([-0.5, -1.2])main_new = torch.tensor([-0.3, -1.0], requires_grad=True)aux_old = torch.tensor([-0.7, -0.9])aux_new = torch.tensor([-0.5, -0.8], requires_grad=True)main_adv = torch.tensor([1.0, -0.5])aux_adv = torch.tensor([-0.3, 0.6])loss, m_loss, a_loss = dapo_loss(main_old, main_new, aux_old, aux_new, main_adv, aux_adv)print(f"DAPO Total Loss: {loss.item():.4f}, Main Loss: {m_loss.item():.4f}, Aux Loss: {a_loss.item():.4f}")## 总结本文深入剖析了 PPO、GRPO、GSPO 和 DAPO 四种策略优化算法的 Loss 计算原理,并提供了可运行的代码示例。PPO 通过裁剪机制保证了策略更新的稳定性;GRPO 引入了群体相对优势,适用于多智能体协作场景;GSPO 使用软约束(如 KL 散度)替代硬裁剪,提供了更灵活的优化框架;DAPO 则通过双智能体耦合损失,处理对抗或协作环境。在实际应用中,选择合适的算法取决于具体问题:对于单智能体任务,PPO 仍是首选;对于多智能体系统,GRPO 和 DAPO 各有侧重;而 GSPO 则适合需要精细控制策略更新幅度的场景。理解这些 Loss 的底层计算,有助于开发者在自定义任务中灵活调整和优化算法。