ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

[Agent Memory / 强化学习] MemPO源码学习笔记 ---(5)--- GRPO

[Agent Memory / 强化学习] MemPO源码学习笔记 ---(5)--- GRPO [Agent Memory / 强化学习] MemPO源码学习笔记 —(5)— GRPO文章目录[Agent Memory / 强化学习] MemPO源码学习笔记 ---(5)--- GRPO0x00 概要0x01 原理1.1 现状1.2 GRPO1.3 PPO vs GRPO1.4 为什么MemPO选GRPO而不是PPOCriticCritic 在多轮长序列中极难训练GRPO 比较适合 outcome-based 稀疏奖励计算资源节约Memory Reward 的特殊性小结0x02 MemPO GRPO2.1 阶段2.2 模型2.3 优势函数2.4 前向传播完整训练步的 Forward Pass 计数特色对比2.5 Loss输入计算ratioImportance sampling张量形状MemPO的特殊之处为什么不用两个独立loss分别优化梯度冲突问题PPO的clip 机制需要统一advantage实现极简计算高效2.6 梯度2.7 KL约束直觉含义直觉类比数学直觉在MemPO中的实际作用实现0xFF 参考0x00 概要现有的基于强化学习的 Memory 管理方法往往缺乏一种有效机制针对 Memory 的更新内容进行引导优化Memory 的内容难以保证质量。MemPO(Self-Memory Policy Optimization)使模型对 Memory 进行自管理并引入了基于有效信息含量的 Memory-level 的优势估计引导 Memory 保留对解决任务更有效的信息进而提升记忆有效性。MemPO的独特切入点让模型把记忆写在每轮开头()形式上像“自我对话的草稿纸“既是记忆又是思考链的一部分。这样变成可训练的策略变量用RL信号端到端地教会模型“什么值得记、怎么记。RL 直接端到端优化这一行为无需额外的记忆模块。MemPO 的信息如下论文标题MemPO: Self-Memory Policy Optimization for Long-Horizon Agents论文地址https://arxiv.org/abs/2603.00680代码地址https://github.com/TheNewBeeKing/MemPO模型和数据集地址https://huggingface.co/collections/NewBeeKing/mempo本篇看看GRPO的使用。0x01 原理1.1 现状Agent 引入记忆机制的目的是通过移除无关信息、保留关键细节来应对智能体的长上下文问题。原始 GRPO 基于答案正确性计算奖励并使用轨迹级别的优势 (advantage)即同一条轨迹内所有 token共享同一个奖励。这导致对记忆生成的奖励信号稀疏、指导有限 — 因为最终答案的正确性无法直接反映交互过程中每次 操作的质量。MemPO设计了一种新颖的优势计算方法在轨迹级优势之外额外评估每一步 中记忆的信息含量并计算一个附加的优势值从而确保记忆在保持简洁的同时保留重要信息。论文原文说的:“computes an additional advantage” 除了 outcome_adv 之外额外再算一个 advantage。 “附加的优势值” 指的就是 Memory Advantage (mem_adv)。对应代码:A2: compute_grpo_memory_advantage () → mem_adv (P_mem - P_full - mean) /std → 仅作用于 … 区间叠加方式(A3):final_advoutcome_advmem_adv ↑ 原有的 ↑附加的(additional)“附加”(additional) 强调的是这是 MemPO 在 GRPO 基础上新增的部分 — 原始 GRPO 只有 outcome_adv, MemPO 额外加了 mem_adv来精确指导 的生成质量。1.2 GRPOGRPO是PPO的一个变体区别仅在于advantage的计算方式(用组内统计替代 Critic)。GRPO 本质上是 用统计方法替代了 Critic网络 — 把一个需要学习的组件 (Critic) 换成了一个不需要学习的统计计算 (组内均值 / 标准差)代价是需要每个question 生成多条轨迹。其他所有环节 (PPO 优化框架) 不受影响。标准PPOadvantageV_critic(s)-R(s)← 需要单独训练一个Critic网络GRPO (Group Relative Policy Optimization):advantage(reward-group_mean)/group_std ← 不需要Critic 其中group同一个question的16条rollout轨迹GRPO 是无Critic的PPO——它保留了PPO 的Clipped surrogate lossImportance sampling ratioKL penalty to ref model多 epoch mini-batch 更新但去掉了 Critic网络用同组轨迹的相对排名代替 value baseline。1.3 PPO vs GRPOGRPO 和 PPO 的核心差异就在于 advantage 的计算方式。其余部分 (clipped loss、ratio、KL、多epoch更新) 完全相同。PPO: adv reward - V (s) ← 需要训练 Critic 来估计 V (s)GRPO: adv (score - mean) /std ← 用同组轨迹统计量替代 V (s)PPO 和 GRPO 对比如下差异汇总如下┌──────────────────┬─────────────────────────────┬─────────────────────────────┐ │ │ 标准 PPO │ GRPO │ ├──────────────────┼─────────────────────────────┼─────────────────────────────┤ │ 轨迹数/question │ 通常1条 │16条(group size)│ ├──────────────────┼─────────────────────────────┼─────────────────────────────┤ │ Critic 网络 │ ✅ 需要(~7B)│ ❌ 不需要 │ ├──────────────────┼─────────────────────────────┼─────────────────────────────┤ │ Advantage 来源 │ GAE:reward-V(s)│(score-mean)/std │ ├──────────────────┼─────────────────────────────┼─────────────────────────────┤ │ 额外训练步骤 │ Critic loss │ 无 │ ├──────────────────┼─────────────────────────────┼─────────────────────────────┤ │ 显存占用 │ actorrefcritic │ actorref │ ├──────────────────┼─────────────────────────────┼─────────────────────────────┤ │ Advantage 精度 │ token-level(但有 │ trajectory-level │ │ │ estimation error)│(无estimation error)│ ├──────────────────┼─────────────────────────────┼─────────────────────────────┤ │ 适合场景 │ dense reward │ sparse/outcome reward │ └──────────────────┴─────────────────────────────┴─────────────────────────────┘1.4 为什么MemPO选GRPO而不是PPOCriticMemPO不用Critic的四个原因如下Critic 在多轮长序列中极难训练标准PPOV(s_t) 需要为序列中每个token位置预测未来累计回报预测“未来能否答对“。但是MemPO的序列结构为[Round1_tokens | Round2_tokens | …| Round5_tokens]。此长度可达数千token奖励仅在最末尾(sparse reward)。因此Critic 面对的挑战序列极长 → 需要巨大容量的value网络奖励极稀疏→ V(s)几乎处处为0难以学到有意义的信号多轮工具交互→状态空间复杂value estimation 噪声大GRPO 比较适合 outcome-based 稀疏奖励GRPO 用同组轨迹均值替代value baseline比较适合trajectory-level离散奖励。GRPO的假设奖励是trajectory-level的标量 → 完美匹配EM check的{0, 1}评分。不需要学习 V(s_t)baseline 同 question 16 条轨迹的均值 (score - mean) / std → 零额外参数零额外训练无 value estimation 误差。计算资源节约PPOCritic:额外一个与actor 同规模的 Critic 网络(7B 参数)Critic需要额外前向反向显存翻倍actor(7B) ref(7B) critic(7B)21B 参数GRPO:仅actor(7B)ref(7B)14B参数省下的资源用于更多并发rollout(16条/question)Memory Reward 的特殊性Memory Reward 本身自带baseline(P_memP_full)无需Critic 估计。mem_rewardP_mem-P_full 这本身就是一个自带baseline的信号如果用 Critic还需要为 区间单独训练 value head → 但的“好坏“取决于未来能否答对(极长时间依赖)→ Critic几乎不可能准确估计这个value。GRPO方案直接跨轨迹归一化mem_reward简单有效小结GRPO在MemPO 场景下是更实用的选择一一稀疏奖励、长序列、多轮交互这三个特点让Critic训练极其困难而 GRPO通过“同组相对排名“巧妙绕过了value estimation问题。MemPO最特色的地方在标准GRPO之上额外为片段设计了细粒度的位置感知奖励让梯度信号可以精确地作用于“记忆写作“行为而不只是笼统地惩奖整条轨迹。我们接下来仔细分析。0x02 MemPO GRPO2.1 阶段GRPO算法 Advantage 计算方式 PPO 优化框架因此具体可以分两个环节环节1GRPO Advantage计算(无梯度)B4-algo:outcome_adv(score-mean)/std ← 纯数值运算 A2:mem_adv(r_t-mean)/std ← 纯数值运算 A3:final_advoutcome_advmem_adv ← 纯加法 所有 advantage 都是detached 常数不参与计算图环节2PPO Update(有梯度)forepochinppo_epochs:formini_batchinshuffie(batch):new_log_probactor.forward(mini_batch)↑梯度计算 ratioexp(new_log_prob-old_log_prob)loss-mean(final_adv x clip(ratio))KL_penalty loss.backward()← 反向传播 optimizer.step()← 模型优化总结GRPO只决定“每个token 该鼓励还是抑制、强度多大”(advantage 值)但“怎么优化模型参数“完全是PPO 的事一一一梯度计算、反向传播、模型更新都在 PPO Update 环节。2.2 模型MemPO 有以下几种模型actor(策略模型)就是正在被训练的LLM(如Qwen2.5-7B)每个PPO step都会更新其参数配置actor_rollout_ref.model.path→初始化自SFT模型既用于rollout生成也用于PPO更新时的前向计算ref_model(参考模型)与actor结构完全相同的 LLM但参数冻结不更新初始化为训练开始时的actor快照(即SFT模型本身)作用计算KL散度惩罚KL(π_actor || π_ref)防止actor偏离初始策略太远(PPO的信任域约束)在代码中的体现run_train.sh 中 actor_rollout_ref.model.pathNewBeeKing/MemPo_Qwen2.5-SFTactor ← 加载这个模型训练中不断更新 ref ← 加载同一个模型训练中冻结 rollout ← 用actor的权重做推理(通过SGLang服务)三者在PPO loss中的角色ratioexp(new_log_prob_actor-old_log_prob_actor)↑当前参数 ↑本轮开始时的快照 KL_penaltyratio_to_ref-log(ratio_to_ref)-1where ratio_to_refexp(log_prob_actor-log_prob_ref)↑永远不更新 loss-adv x clip(ratio)KL_coef x KL_penalty简单来说actor 学生不断学习改进ref_model “老师基线”, 确保学生不会偏离太远old_log_prob “上一次考试成绩”, 用于计算 importance sampling ratio2.3 优势函数Outcome Advantage 和 Memory Advantage 两者都用 GRPO 风格的归一化方式计算 advantage但侧重点不同。Outcome Advantage ——— GRPO 标准流程B4-algo: compute_grpo_outcome_advantage()分组同一 question 的 16 条轨迹归一化adv (score - group_mean) / group_std→ 这就是 GRPO 的核心—用组内相对排名替代 CriticMemory Advantage ——— GRPO 风格但维度不同A2: compute_grpo_memory_advantage()分组同一 question 的所有轨迹 × 所有轮次 (~48 个值)归一化adv (mem_reward - pool_mean) / pool_std→ 借鉴了 GRPO 的 “组内归一化” 思想→ 但池化范围更大(跨轨迹 跨轮次)两者最终final_adv outcome_adv mem_adv → 送入同一个 PPO loss严格来说Outcome Advantage 标准 GRPOMemory Advantage GRPO 启发的归一化(不是 GRPO 论文中定义的是 MemPO 的创新设计)最终优化 PPO clipped surrogate loss(GRPO 只是 advantage 计算方式优化器仍是 PPO)2.4 前向传播此点在Rollout篇也有涉及。完整训练步的 Forward Pass 计数MemPO 相比原版GRPO 多了1次extra forward pass(步骤②)但该次同时批量处理了 full_traj 和mem_traj实际吞吐开销约是标准old_log_prob的1.5~2倍是 MemPO最主要的训练额外成本。─────────────────────────────────────────────────── ① 生成阶段(generate_sequences)SGLang 自回归解码共 n16条轨迹 → 本质也是 forward但 KV cache 优化计一次 ─────────────────────────────────────────────────── ②★ MemPO 专属:compute_log_prob(full_trajmem_traj)agent_loop.py concat[全部 full_traj,全部 mem_traj]→ 一次调用但序列数量2× B ×(T-1)× n Bbatch_size,T轮次,n16─────────────────────────────────────────────────── ③ compute_log_prob(old_log_prob)ray_trainer.py → actor 计算轨迹的旧 logp(供 PPO ratio 使用)─────────────────────────────────────────────────── ④ compute_ref_log_prob(KL 约束)ray_trainer.py → ref model 计算 logp(供 KL 惩罚使用)─────────────────────────────────────────────────── ⑤ actor update(多 epoch 反向传播)默认 ppo_epochs1,每次需要当前 logp特色特性详情是否推理两次?否1次extra forward pass不生成新 token实际操作对已生成的答案 Z用两种不同的输入上下文计算 logp计算次数一次 forward pass两种输入拼成一个 batch计算时机rollout 完成后advantage 计算前目的衡量“仅凭能否预测正确答案“的能力对比与原版 VeRL的对比阶段生成原版 VeRL GRPOMemPO生成①generate①generate记忆奖励✗无✓②fullmem 双路 logp旧logp③old_log_prob③old_log_probref logp④ref_log_prob④ref_log_prob更新⑤actor update⑤actor update总计4次 forward5次forward每个 search_results 都是一次搜索的返回不是多次搜索的集合。2.5 Lossoutcome_adv 和 mem_adv 两者共同作为 PPO 的 advantage 信号在同一个 PPO loss 中训练。final_advoutcome_advmem_adv ← 叠加后作为 PPO 的 advantage PPO loss-mean(final_adv × clip(ratio,1-ε,1ε)× response_mask)KL_coef × KL(π||π_ref)不是两个独立的训练过程而是: 一次前向 → 一个 loss → 一次反向传播不同 token 接收到的 advantage 值不同:position:[R1 tokens][memR2 tokens/mem][think tokens][memR3/mem][answer]final_adv:[0.8][0.81.2][0.8][0.8-0.5][0.8]↑ 仅 outcome ↑ outcomemem(正)↑ 仅 outcome ↑ outcomemem(负)输入loss 的输入如下new_log_prob[bsz,seq_len]← 当前actor前向得到 old_log_prob[bsz,seq_len]← rollout时的快照(detached)ref_log_prob[bsz,seq_len]← refmodel(冻结)final_adv[bsz,seq_len]← outcome_advmem_adv response_mask[bsz,seq_len]←1response token0prompt token计算计算公式为PPO loss -mean( final_adv × clip(ratio, 1-ε, 1ε) × response_mask ) KL_coef × KL(π || π_ref)PPO是REINFORCE的改进版加入importance sampling ratioratioπ_new/π_old允许在旧数据上多次更新加入clip约束防止ratio偏离太大(限制单步更新幅度)加入KL penalty防止偏离参考策略太远本质上PPO loss中的final_adv × ratio就是REINFORCE梯度的importance-weighted版本。# Importance Sampling Ratioratioexp(new_log_prob-old_log_prob)# Clipped Surrogatesurr1ratio × final_adv surr2clip(ratio,1-e,1e)× final_adv policy_loss-mean(min(surr1,surr2)× response_mask)# KL Penalty (low-variance estimator)ratio_refexp(new_log_prob-ref_log_prob)klratio_ref-log(ratio_ref)-1kl_lossKL_coef × mean(kl × response_mask)# Totaltotal_losspolicy_losskl_loss最终的 total_loss 是一个标量(scalar)但中间步骤涉及张量运算即total_loss 是标量。中间的 ratio、advantage等是[bszseq_len]张量通过mean()操作压缩为标量。total_loss.backward()从该标量反传梯度到所有模型参数唯一有梯度的量是new_log_prob(当前actor前向计算得到)。我们接下来看看几个具体细节。ratioPPO loss中的ratio是importance sampling比率配合clip防止单步更新过大ratio 1当前策略比生成时更倾向选择这些tokenratio 1当前策略比生成时更不倾向选择这些tokenratio 1没变化Importance samplingImportance sampling 允许用一个分布旧策略π_old采的样本来估计另一个分布新策略π_new下的期望E_π_new[f(x)]E_π_old[f(x)× π_new(x)/π_old(x)]PPO需要它是因为rollout在旧策略下生成但要更新到新策略。ratioπ_new/π_old就是importanceweight纠正了分布不匹配。clip限制ratio范围是为了防止weight太大导致高方差。张量形状各步骤的形状如下new_log_prob[bsz,seq_len]← 张量(每个token一个log概率)old_log_prob[bsz,seq_len]← 张量 ratio[bsz,seq_len]← 张量(逐token计算)final_adv[bsz,seq_len]← 张量(逐token不同值)surr1[bsz,seq_len]← 张量 surr2[bsz,seq_len]← 张量min(surr1,surr2)[bsz,seq_len]← 张量 response_mask[bsz,seq_len]← 张量(0/1)policy_loss-mean(min(...)× mask)← 标量 ✓(对所有元素求平均)kl_lossKL_coef × mean(...)← 标量 ✓ total_losspolicy_losskl_loss ← 标量 ✓ total_loss.backward()← 从这个标量反传梯度到所有参数关键mean()操作将[bszseq_len的张量压缩为标量然后.backward()从该标量计算梯度。MemPO的特殊之处Loss公式本身没有修改特殊性全在final_adv的构造上标准 GRPO:final_adv[i,:]outcome_adv_i ← 全序列同一个值 MemPO:final_adv[i,:]outcome_adv_imem_adv[i,:]↑ 仅mem区间非零效果同一条轨迹内不同 token 的 advantage 值不同token类型 advantage 值 梯度效果 ─────────────────────────────────────────────────────────mem.../memoutcomemem_adv_t 双重驱动(可正可负)think内容 outcome 仅结果驱动searchquery outcome 仅结果驱动answer内容 outcome 仅结果驱动 prompt tokens masked out(0)无梯度这就是 MemPO 的全部创新点在 loss 层面的体现——通过让 token 接收额外的 memory quality信号实现对记忆摘要能力的精准优化而不影响其他 token 的学习。为什么不用两个独立loss分别优化MemPO 的训练本质上是一个 PPO 算法 一个复合 advantage:GRPO (Group Relative Policy Optimization) 负责计算 outcome_advMemory reward 机制负责计算 mem_adv两者加和后送入标准 PPO clipped surrogate loss并没有分别训练两个 objective, 也没有分步交替优化 — — — 就是一个统一的梯度更新。如此设计的原因有三梯度冲突问题如果两个loss独立loss_outcome-outcome_adv × log π(token)loss_memory-mem_adv × log π(mem_token)以场景3(答错好mem)为例:loss_outcome 想让memtoken概率 ↓(因为整条轨迹答错了)loss_memory 想让memtoken概率 ↑(因为摘要写得好)梯度冲突两个loss对token可能方向相反 → 两个梯度方向相反 → 训练不稳定、震荡。叠加方案直接解决(-0.8) (1.2) 0.4产出一个明确的净方向。PPO的clip 机制需要统一advantagePPO clip的含义限制每步更新幅度如果拆成两个loss分别clip每个loss各允许幅度的更新叠加后实际更新了2ε→超出信任域统一后只clip一次final_adv outcome mem → clip一次 → 总更新在ε内实现极简计算高效统一方案1次前向1次反向1次参数更新 → 实现极简计算高效 → 无需调两个1oss的权重系数final_advoutcome_advmem_adv ←1行加法独立方案→2次前向2次反向2次参数更新(或需要梯度累积) → 还需要调两个loss的权重系数 →实际上调权重系数 ≈ 调mem_adv的相对幅度(本质相同)总结叠加 advantage 本质上等价于带权重的多目标优化但更稳定、更高效、且天然兼容PPO的clip约束2.6 梯度我们来看看整个训练过程中哪些计算有梯度、哪些没有。✗ 无梯度(detached/frozen)✗ 无梯度(detached/frozen)-Rollout 生成 token → SGLang 推理不保留计算图-A1 compute_log_prob(mem_reward)→ detached仅算数值-B2 compute_score(outcome_reward)→ 纯字符串匹配无张量-B4-algo outcome_adv → 纯数值运算-A2 mem_adv → 纯数值运算-A3 final_advoutcomemem → 常数张量-old_log_prob → detached 快照-ref_log_prob → 冻结模型☑ 有梯度(唯一来源)☑ 有梯度(唯一来源) PPO Update 中 new_log_probactor.forward(response_ids)← 当前 actor 前向 ↑ 这是唯一参与计算图的量 ratioexp(new_log_prob-old_log_prob)← 梯度流经 new_log_prob loss-mean(final_adv × clip(ratio)× mask)← 标量KL_coef x f(new_log_prob,ref_log_prob)loss.backward()↓ 梯度方向 loss → ratio → new_log_prob → actor 参数(weights,biases,embeddings)optimizer.step()→更新actor 所有参数因此整个 MemPO 的梯度就是∂loss/∂θ_actor(PPO Update 中 new_log_prob 对 actor 参数的梯度)通过 new_log_prob这一个计算图节点反传到模型参数。所有advantage、reward、ref/old log_prob 都只是常数系数决定梯度的方向和大小”但不贡献梯度本身。2.7 KL约束直觉含义KL约束是防止模型为了刷分而跑偏一确保更新后的策略不会偏离起点太远。直觉类比想象一个学生(actor)在刷题提分没有约束可能发现某种作弊捷径(如固定输出某个高频答案)→分数暂时上升但能力退化KL约束“你可以改进但不能变得和原来的自己(ref)差别太大” →保持泛化能力的同时逐步提升数学直觉KL(π_actor || π_ref)衡量两个分布的距离 0actor和ref完全一样(没学到任何东西)很大actor和ref差别巨大(可能过度优化/rewardhacking)在loss中total_losspolicy_lossKL_coef x KL policy_loss想让模型往高reward方向走→ 拉离ref KL_penalty想让模型别离ref太远→ 拉回ref两者对抗 → 模型在提升和稳定之间找到平衡在MemPO中的实际作用KL penalty约束actor不偏离ref model(SFT起点)太远。没有KL penalty时可能出现的问题模型学会写一种固定模板的(高mem_reward但无实际信息 / reward hacking)模型对所有问题都搜同一个query(碰巧某些场景有效)输出多样性崩溃(所有16条轨迹趋同)训练不稳定KL penalty 确保→ 模型的输出分布保持多样性 → 每步更新幅度有限训练稳定 → 不会出现reward hacking实现配置kl_loss_type low_var_kl使用 k3估计量(Schulman 2020)# 普通KL(k1有偏梯度)kl ≈ log π_ref-log π_new# Low-variance KL estimator# 代码中用的不是直接的 KL而是 low-variance 近似# low_var_kl(k3无偏梯度)klratio_ref-log(ratio_ref)-1where ratio_refexp(new_log_prob-ref_log_prob)π_actor/π_ref 当 ratio_ref1(完全一样):kl1-0-10✅ 当 ratio_ref1(actor概率更高):kl0当 ratio_ref1(actor概率更低):kl0→ 任何偏离都被惩罚是一个弹簧力把 actor 拉回 ref为什么用 k3 k3比 k1方差更低训练更稳定尤其适合多轮 agent 场景(轨迹长、方差本来就大)。几个估计器比较如下估计器公式性质k1log ⁡ π θ − log ⁡ π ref \log\pi_\theta -\log\pi_{\text{ref}}logπθ​−logπref​无偏方差大可负k21 2 ( log ⁡ π θ log ⁡ π ref ) 2 \frac{1}{2}(\log\pi_\theta\log\pi_{\text{ref}})^221​(logπθ​logπref​)2有偏恒正k3r − 1 − log ⁡ r r -1 -\log rr−1−logr(r π ref / π θ r\pi_{\text{ref}}/\pi_\thetarπref​/πθ​)无偏 恒正 低方差(Schulman)low_var_kl同k3再做clamp(-10, 10)数值保护同k3 防爆0xFF 参考
返回列表