发布时间:2026/7/25 0:20:45
PPO GRPO GSPO DAPO的Loss计算与代码实现 PPO、GRPO、GSPO、DAPO 的 Loss 计算与代码实现在强化学习Reinforcement Learning, RL领域策略优化算法一直是研究的核心。从经典的 PPOProximal Policy Optimization到近年来出现的 GRPOGroup Relative Policy Optimization、GSPOGeneralized Surrogate Policy Optimization以及 DAPODual-Agent Policy Optimization这些算法通过不同的 Loss 设计解决了策略更新中的稳定性、样本效率以及多智能体协作等问题。本文将深入剖析这四种算法的 Loss 计算原理并提供可运行的代码片段帮助读者从底层理解其工作机制。## PPO基于信任区域的策略优化PPOProximal 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, epsilon0.2, entropy_coef0.01, value_coef0.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_gradFalse)new_log_probs torch.tensor([-0.3, -1.0, -0.6], requires_gradTrue)advantages torch.tensor([1.0, -0.5, 0.8])values torch.tensor([0.9, 0.3, 0.7], requires_gradTrue)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(fPPO Total Loss: {loss.item():.4f}, Policy Loss: {p_loss.item():.4f}, Value Loss: {v_loss.item():.4f})## GRPO群体相对策略优化GRPOGroup 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{j1}^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, epsilon0.2, beta0.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(dim1, keepdimTrue).expand_as(old_log_probs) kl_div F.kl_div(new_log_probs, group_log_probs, reductionbatchmean, log_targetTrue) # 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_gradTrue)rewards torch.tensor([[1.0, 0.5], [0.8, 1.2], [0.3, 0.7]])group_rewards rewards.mean(dim1, keepdimTrue) # 群体平均loss, p_loss, kl grpo_loss(old_log_probs, new_log_probs, rewards, group_rewards)print(fGRPO Total Loss: {loss.item():.4f}, Policy Loss: {p_loss.item():.4f}, KL Div: {kl.item():.4f})## GSPO广义替代策略优化GSPOGeneralized 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_coef0.1, distancekl): 计算 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, reductionbatchmean, log_targetTrue) kl2 F.kl_div(old_log_probs, m_log_probs, reductionbatchmean, log_targetTrue) 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_gradTrue)advantages torch.tensor([1.0, -0.5, 0.8])loss, surr, kl gspo_loss(old_log_probs, new_log_probs, advantages, lambda_coef0.5, distancekl)print(fGSPO Total Loss: {loss.item():.4f}, Surrogate: {surr.item():.4f}, KL: {kl.item():.4f})## DAPO双智能体策略优化DAPODual-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, alpha0.5, epsilon0.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_gradTrue)aux_old torch.tensor([-0.7, -0.9])aux_new torch.tensor([-0.5, -0.8], requires_gradTrue)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(fDAPO 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 的底层计算有助于开发者在自定义任务中灵活调整和优化算法。

相关新闻

2026/7/25 0:20:45

阿里云智能语音简单使用:语音识别

阿里云智能语音简单使用:语音识别 1. 什么是阿里云智能语音识别?阿里云智能语音识别(ASR,Automatic Speech Recognition)是阿里云提供的一项人工智能服务,能够将音频中的语音实时或离线转换成文字。这项技术…

2026/7/25 0:20:45

浏览器端EPUB构建技术栈:零部署的现代电子书编辑解决方案

浏览器端EPUB构建技术栈:零部署的现代电子书编辑解决方案 【免费下载链接】EPubBuilder 一款在线的epub格式书籍编辑器 项目地址: https://gitcode.com/gh_mirrors/ep/EPubBuilder 技术挑战:传统电子书编辑器的架构困境 在数字化内容创作领域&am…

2026/7/25 0:20:45

如何快速定位Windows热键冲突:完整的热键侦探指南

如何快速定位Windows热键冲突:完整的热键侦探指南 【免费下载链接】hotkey-detective A small program for investigating stolen key combinations under Windows 7 and later. 项目地址: https://gitcode.com/gh_mirrors/ho/hotkey-detective 你是否曾遇到…

2026/7/25 1:45:50

VMware虚拟机中启用Windows 7 Build 7106 Aero特效的完整指南

在 Windows 7 的早期开发阶段,Build 7106 是一个具有里程碑意义的版本,它首次向公众展示了完整的 Aero 桌面特效。对于技术爱好者、怀旧系统研究者或需要特定测试环境的开发者而言,在虚拟机中复现这一经典环境,不仅是对历史的回顾,更是一次深入理解操作系统图形子系统、驱…

2026/7/25 1:45:50

2026年Linux运维/SRE学习路线:从零基础到云原生与AIOps实战

如果你在2026年还在用2018年的方法学Linux运维,那你可能已经落后了整整一个技术代际。 这不是危言耸听。过去几年,运维领域经历了从“手工救火”到“平台工程”,再到如今“AI驱动”的深刻变革。传统的“命令大全”式学习路径,面对云原生、可观测性、SRE工程实践和AIOps的复…

2026/7/25 1:45:50

Linux运维/SRE零基础到求职:2024年系统学习路径与实战指南

你是不是也刷到过那些“零基础三个月转行运维,月入过万”的广告?是不是也收藏了一堆“Linux命令大全”、“SRE面试宝典”,但打开后面对海量碎片化信息,依然不知道从何下手,更不清楚学完这些到底能不能找到工作? 这正是当前想进入Linux运维/SRE领域的新人最真实的困境: …

2026/7/25 1:40:50

C++实现RSA模幂运算:重复平方乘算法详解与优化

1. 项目概述:当RSA遇上大数幂运算如果你尝试过用C手搓一个RSA加密算法,或者仅仅是好奇想实现一下,那么“大数幂”这个计算绝对是你绕不过去的一道坎。想象一下,你要计算一个像123456789^987654321 mod 1000000007这样的表达式。直…

2026/7/23 12:54:51

Unity与Python本地通信:基于Flask的跨语言数据交换实战

1. 项目概述:为什么我们需要一个本地通信服务器?在游戏开发、数字孪生、仿真训练等众多领域,Unity作为强大的实时3D内容创作平台,其核心逻辑通常由C#驱动。然而,当我们需要进行复杂的数据分析、机器学习推理、科学计算…

2026/7/25 0:00:15

C++ string类模拟实现:从深拷贝到内存管理的完整指南

1. 项目概述:为什么我们要“手撕”string类?在C的学习道路上,尤其是从C语言过渡到C的“初阶”阶段,string类绝对是一个绕不开的核心。标准库里的std::string用起来太方便了,、find、substr,几个操作符和函数…

2026/7/25 0:00:15

三角洲寻宝鼠工具:高效文件搜索与资源管理实战指南

1. 先搞清楚“三角洲寻宝鼠”到底是什么工具从名称来看,“三角洲寻宝鼠”更像是一个资源查找或文件检索类工具,而不是游戏或娱乐软件。这类工具的核心价值在于帮助用户快速定位特定资源,比如文档、图片、压缩包或特定格式的文件。如果你经常需…

2026/7/25 0:00:15

VHF 甚高频语音喊话系统(桥梁智能防撞场景)核心优势

一、直达船员,预警链路最短营运船舶强制标配 VHF 船载电台,属于驾驶室常态化值守设备;预警语音直接传递至驾驶人员,区别于岸上声光报警(船员经常听不到)、短信 / 小程序(船员极少主动查看&#…

2026/7/25 0:59:36

3个高效策略:快速掌握Axure中文界面配置

3个高效策略:快速掌握Axure中文界面配置 【免费下载链接】axure-cn Chinese language file for Axure RP. Axure RP 简体中文语言包。支持 Axure 11、10、9。不定期更新。 项目地址: https://gitcode.com/gh_mirrors/ax/axure-cn 还在为Axure RP的英文界面感…