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

文章来源声明: 原文作者:罗西的思考; 来源站点:掘金; 原文链接:https://juejin.cn/post/7690413193491939328; 本文基于上述来源整理/加工,觅优补充点评,仅供技术学习交流。版权归原作者所有。
觅优短评

为长周期Agent的记忆管理提供了高效的RL优化思路,绕过Critic训练难题,适用于稀疏奖励、多轮交互的复杂任务场景。

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

0x00 概要

现有的基于强化学习的 Memory 管理方法往往缺乏一种有效机制针对 Memory 的更新内容进行引导优化,Memory 的内容难以保证质量。

MemPO(Self-Memory Policy Optimization)使模型对 Memory 进行自管理,并引入了基于有效信息含量的 Memory-level 的优势估计,引导 Memory 保留对解决任务更有效的信息,进而提升记忆有效性。

MemPO的独特切入点:让模型把记忆写在每轮开头(),形式上像“自我对话的草稿纸“,既是记忆又是思考链的一部分。这样,变成可训练的策略变量,用RL信号端到端地教会模型“什么值得记、怎么记"。RL 直接端到端优化这一行为,无需额外的记忆模块。

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):<span>final_adv</span> = outcome_adv + mem_adv
                          ↑ 原有的        ↑ "附加的"(additional)

"附加"(additional) 强调的是:这是 MemPO 在 GRPO 基础上新增的部分 — 原始 GRPO 只有 outcome_adv, MemPO 额外加了 mem_adv来精确指导 的生成质量。

1.2 GRPO

GRPO是PPO的一个变体,区别仅在于advantage的计算方式(用组内统计替代 Critic)。GRPO 本质上是 " 用统计方法替代了 Critic网络 "— 把一个需要学习的组件 (Critic) 换成了一个不需要学习的统计计算 (组内均值 / 标准差),代价是需要每个question 生成多条轨迹。其他所有环节 (PPO 优化框架) 不受影响。

标准PPO:

advantage = <span>V_critic</span>(s) - <span>R</span>(s)    ←     需要单独训练一个Critic网络

GRPO (Group Relative Policy Optimization):

<span>advantage</span> =(reward-group_mean)/ group_std  ←   不需要Critic 
其中<span>group</span> = 同一个question的<span>16</span>条rollout轨迹

GRPO 是"无Critic的PPO"——它保留了PPO 的:

  • Clipped surrogate loss
  • Importance sampling ratio
  • KL penalty to ref model
  • 多 epoch mini-batch 更新

但去掉了 Critic网络,用同组轨迹的相对排名代替 value baseline。

1.3 PPO vs GRPO

GRPO 和 PPO 的核心差异就在于 advantage 的计算方式。其余部分 (clipped loss、ratio、KL、多epoch更新) 完全相同。

  • PPO: adv = reward - V (s) ← 需要训练 Critic 来估计 V (s)
  • GRPO: adv = (score - mean) /std ← 用同组轨迹统计量替代 V (s)

PPO 和 GRPO 对比如下:

5-PPO 和 GRPO 对比

5-PPO 和 GRPO 对比

差异汇总如下:

┌──────────────────┬─────────────────────────────┬─────────────────────────────┐
│                  │ 标准 PPO                     │ GRPO                        │
├──────────────────┼─────────────────────────────┼─────────────────────────────┤
│ 轨迹数/question   │ 通常<span>1</span>条                     │ <span>16</span>条(group size)            │
├──────────────────┼─────────────────────────────┼─────────────────────────────┤
│ Critic 网络       │ ✅ 需要(~<span>7</span>B)               │ ❌ 不需要                    │
├──────────────────┼─────────────────────────────┼─────────────────────────────┤
│ Advantage 来源    │ GAE: <span>reward-V</span>(s)            │ (score-mean)/std            │
├──────────────────┼─────────────────────────────┼─────────────────────────────┤
│ 额外训练步骤      │ Critic loss                  │ 无                           │
├──────────────────┼─────────────────────────────┼─────────────────────────────┤
│ 显存占用          │ actor+ref+critic             │ actor+ref                  │
├──────────────────┼─────────────────────────────┼─────────────────────────────┤
│ Advantage 精度    │ <span>token-level</span>(但有             │ trajectory-level          │
│                  │ estimation error)            │ (无estimation error)     │
├──────────────────┼─────────────────────────────┼─────────────────────────────┤
│ 适合场景          │ dense reward                 │ sparse/outcome reward       │
└──────────────────┴─────────────────────────────┴─────────────────────────────┘

1.4 为什么MemPO选GRPO而不是PPO+Critic

MemPO不用Critic的四个原因如下:

Critic 在多轮长序列中极难训练

标准PPO:V(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 误差。
计算资源节约

PPO+Critic:

  • 额外一个与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_mem-P_full),无需Critic 估计。

<span>mem_reward</span>=P_mem-P_full     这本身就是一个<span>"自带baseline"</span>的信号

如果用 Critic,还需要为 区间单独训练 value head → 但的“好坏“取决于未来能否答对(极长时间依赖)→ Critic几乎不可能准确估计这个value。

GRPO方案:直接跨轨迹归一化mem_reward,简单有效

小结

GRPO在MemPO 场景下是更实用的选择一一稀疏奖励、长序列、多轮交互这三个特点让Critic训练极其困难,而 GRPO通过“同组相对排名“巧妙绕过了value estimation问题。

MemPO最特色的地方:在标准GRPO之上,额外为片段设计了细粒度的位置感知奖励,让梯度信号可以精确地作用于“记忆写作“行为,而不只是笼统地惩奖整条轨迹。

我们接下来仔细分析。

0x02 MemPO GRPO

2.1 阶段

GRPO算法 = Advantage 计算方式 + PPO 优化框架,因此具体可以分两个环节:

环节1:GRPO Advantage计算(无梯度)

B4-algo: <span>outcome_adv</span> =(score - mean) / std     ← 纯数值运算
A2: <span>mem_adv</span> = (r_t-mean) / std                 ← 纯数值运算
A3: <span>final_adv</span> = outcome_adv + mem_adv          ← 纯加法
所有 advantage 都是detached 常数,不参与计算图

环节2:PPO Update(有梯度)

for epoch in ppo_epochs:
    for mini_batch in <span>shuffie</span>(batch):
        new_log_prob = actor.<span>forward</span>(mini_batch)
                               ↑梯度计算
                               
        ratio = <span>exp</span>(new_log_prob - old_log_prob) 
        loss = <span>-mean</span>(final_adv x <span>clip</span>(ratio)) + KL_penalty
        
        loss.<span>backward</span>()   ← 反向传播 
        optimizer.<span>step</span>()  ← 模型优化

总结: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 中:
    <span>actor_rollout_ref.model.path</span>=<span>"NewBeeKing/MemPo_Qwen2.5-SFT"</span>
    
    actor   ← 加载这个模型,训练中不断更新
    ref     ← 加载同一个模型,训练中冻结
    rollout ← 用actor的权重做推理(通过SGLang服务)

三者在PPO loss中的角色:

<span>ratio</span> = exp(new_log_prob_actor - old_log_prob_actor)
               ↑当前参数              ↑本轮开始时的快照
<span>KL_penalty</span> = ratio_to_ref -log(ratio_to_ref)-<span>1</span>
       where <span>ratio_to_ref</span> = exp(log_prob_actor -log_prob_ref)
                                                          ↑永远不更新
<span>loss</span> =-adv x clip(ratio)+ KL_coef x KL_penalty 

简单来说:

  • actor = "学生,不断学习改进

  • ref_model = "老师基线", 确保学生不会偏离太远

  • old_log_prob = "上一次考试成绩", 用于计算 importance sampling ratio

2.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 的核心—用组内相对排名替代 Critic

Memory 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 = 标准 GRPO
  • Memory 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 自回归解码,共 <span>n</span>=<span>16</span> 条轨迹
    → 本质也是 forward,但 KV cache 优化,计一次

───────────────────────────────────────────────────
②★ MemPO 专属: compute_log_prob(full_traj + mem_traj)
    agent_loop.py 
    <span>concat</span> = [全部 full_traj, 全部 mem_traj]
    → 一次调用,但序列数量 = 2 × B × (T-1) × n
    <span>B</span>=batch_size, T=轮次, n=<span>16</span>

───────────────────────────────────────────────────
③ 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 反向传播)
 默认 <span>ppo_epochs</span>=<span>1</span>, 每次需要当前 logp

特色
特性详情
是否"推理两次"?否,1次extra forward pass,不生成新 token
实际操作对已生成的答案 Z,用两种不同的输入上下文计算 logp
计算次数一次 forward pass,两种输入拼成一个 batch
计算时机rollout 完成后,advantage 计算前
目的衡量“仅凭能否预测正确答案“的能力
对比

与原版 VeRL的对比

阶段生成原版 VeRL GRPOMemPO
生成①generate①generate
记忆奖励✗无✓②full+mem 双路 logp
旧logp③old\_log\_prob③old\_log\_prob
ref logp④ref\_log\_prob④ref\_log\_prob
更新⑤actor update⑤actor update
总计4次 forward5次forward

每个 search_results 都是一次搜索的返回,不是多次搜索的集合。

2.5 Loss

outcome_adv 和 mem_adv 两者共同作为 PPO 的 advantage 信号,在同一个 PPO loss 中训练。

<span>final_adv</span> = outcome_adv + mem_adv    ← 叠加后作为 PPO 的 advantage

PPO <span>loss</span> = -mean( final_adv × clip(ratio, <span>1</span>-ε, <span>1</span>+ε) × response_mask ) + KL_coef × KL(π || π_ref)

不是两个独立的训练过程,而是: 一次前向 → 一个 loss → 一次反向传播

不同 token 接收到的 advantage 值不同:

position:    <span>[R1 tokens]</span> <span>[<mem>R2 tokens</mem>]</span> <span>[think tokens]</span> <span>[<mem>R3</mem>]</span> <span>[answer]</span>
final_adv:   <span>[  +0.8  ]</span>    <span>[ +0.8 + 1.2 ]</span>      <span>[  +0.8  ]</span>      <span>[+0.8 - 0.5 ]</span>   <span>[ +0.8 ]</span>
              ↑ 仅 outcome  ↑ outcome + mem (正)  ↑ 仅 outcome     ↑ outcome + mem (负)

输入

loss 的输入如下:

new_log_prob <span>[bsz, seq_len]</span>   ←  当前actor前向得到
old_log_prob <span>[bsz, seq_len]</span>   ←  rollout时的快照(detached)
ref_log_prob <span>[bsz, seq_len]</span>   ←  <span>refmodel</span>(冻结)
final_adv    <span>[bsz, seq_len]</span>   ←  outcome_adv + mem_adv
response_mask <span>[bsz, seq_len]</span>  ←  <span>1</span> = response token,<span>0</span> = prompt token

计算

计算公式为:

PPO loss = <span>-mean</span>( final_adv × clip(ratio, <span>1</span>-ε, <span>1</span>+ε) × response_mask ) + KL_coef × <span>KL</span>(π || π_ref)

PPO是REINFORCE的改进版:

  • 加入importance sampling ratio:ratio=π_new/π_old,允许在旧数据上多次更新
  • 加入clip约束:防止ratio偏离太大(限制单步更新幅度)
  • 加入KL penalty:防止偏离参考策略太远
  • 本质上PPO loss中的final_adv × ratio就是REINFORCE梯度的importance-weighted版本。
<span># Importance Sampling Ratio</span>
<span>ratio</span> = exp(new_log_prob - old_log_prob)

<span># Clipped Surrogate</span>
<span>surr1</span> = ratio × final_adv
<span>surr2</span> = clip(ratio,<span>1</span>-e,<span>1</span>+e) × final_adv
<span>policy_loss</span> = -mean( min(surr1, surr2) × response_mask )

<span># KL Penalty (low-variance estimator)</span>
<span>ratio_ref</span> = exp(new_log_prob - ref_log_prob) 
<span>kl</span> = ratio_ref-log(ratio_ref)-<span>1</span>
<span>kl_loss</span> = KL_coef × mean(kl × response_mask)

<span># Total</span>
<span>total_loss</span> = policy_loss + kl_loss 

最终的 total_loss 是一个标量(scalar),但中间步骤涉及张量运算,即total_loss 是标量。中间的 ratio、advantage等是[bsz,seq_len]张量,通过mean()操作压缩为标量。total_loss.backward()从该标量反传梯度到所有模型参数,唯一有梯度的量是new_log_prob(当前actor前向计算得到)。

我们接下来看看几个具体细节。

ratio

PPO loss中的ratio是importance sampling比率,配合clip防止单步更新过大:

  • ratio > 1:当前策略比生成时,更倾向选择这些token
  • ratio < 1:当前策略比生成时,更不倾向选择这些token
  • ratio = 1:没变化
Importance sampling

Importance sampling 允许用一个分布(旧策略 π_old)采的样本来估计另一个分布(新策略 π_new)下的期望:

E_π_new<span>[f(x)]</span> = E_π_old<span>[f(x) × π_new(x) / π_old(x)]</span>

PPO需要它是因为rollout在旧策略下生成,但要更新到新策略。ratio=π_new/π_old就是importance weight,纠正了分布不匹配。clip限制ratio范围是为了防止weight太大导致高方差。

张量形状

各步骤的形状如下:

  new_log_prob     <span>[bsz, seq_len]</span>    ← 张量(每个token一个log概率)
  old_log_prob     <span>[bsz, seq_len]</span>    ← 张量
  ratio            <span>[bsz, seq_len]</span>    ← 张量(逐token计算)
  final_adv        <span>[bsz, seq_len]</span>    ← 张量(逐token不同值)
  surr1            <span>[bsz, seq_len]</span>    ← 张量
  surr2            <span>[bsz, seq_len]</span>    ← 张量
  <span>min</span>(surr1,surr2) <span>[bsz, seq_len]</span>    ← 张量
  response_mask    <span>[bsz, seq_len]</span>    ← 张量(<span>0</span>/<span>1</span>)
  
  policy_loss = <span>-mean</span>(min(...) × <span>mask</span>)  ← 标量 ✓  (对所有元素求平均)
  kl_loss = KL_coef × <span>mean</span>(...)         ← 标量 ✓
  total_loss = policy_loss + kl_loss    ← 标量 ✓

total_loss<span>.backward</span>() ← 从这个标量反传梯度到所有参数

关键:mean()操作将[bsz,seq_len]的张量压缩为标量,然后.backward()从该标量计算梯度。

MemPO的特殊之处

Loss公式本身没有修改,特殊性全在final_adv的构造上:

标准 GRPO:
    final_adv<span>[i, :]</span> = outcome_adv_i   ← 全序列同一个值

MemPO:
    final_adv<span>[i, :]</span> = outcome_adv_i + mem_adv<span>[i, :]</span>
                                        ↑ 仅 <mem> 区间非零

效果:同一条轨迹内,不同 token 的 advantage 值不同:

token类型            advantage 值            梯度效果
─────────────────────────────────────────────────────────
<span><<span>mem</span>></span>...<span></<span>mem</span>></span>      outcome + mem_adv_t      双重驱动(可正可负)
<span><<span>think</span>></span>内容          outcome                  仅结果驱动
<span><<span>search</span>></span>query       outcome                  仅结果驱动
<span><<span>answer</span>></span>内容         outcome                  仅结果驱动
prompt tokens       masked out (=0)          无梯度

这就是 MemPO 的全部创新点在 loss 层面的体现——通过让 token 接收额外的 memory quality 信号,实现对记忆摘要能力的精准优化,而不影响其他 token 的学习。

为什么不用两个独立loss分别优化?

MemPO 的训练,本质上是一个 PPO 算法 + 一个复合 advantage:

  • GRPO (Group Relative Policy Optimization) 负责计算 outcome_adv
  • Memory reward 机制负责计算 mem_adv
  • 两者加和后送入标准 PPO clipped surrogate loss

并没有分别训练两个 objective, 也没有分步交替优化 — — — 就是一个统一的梯度更新。

如此设计的原因有三:

梯度冲突问题

如果两个loss独立:

<span>loss_outcome</span> = -outcome_adv × log π(token) 
<span>loss_memory</span> = -mem_adv × log π(mem_token)

以场景3(答错+好mem)为例:

loss_outcome 想让<mem>token概率 ↓ (因为整条轨迹答错了)
loss_memory  想让<mem>token概率 ↑ (因为摘要写得好)

梯度冲突:两个loss对token可能方向相反 → 两个梯度方向相反 → 训练不稳定、震荡。

叠加方案直接解决:(-0.8) + (+1.2) = +0.4,产出一个明确的净方向。

PPO的clip 机制需要统一advantage

PPO clip的含义:限制每步更新幅度

如果拆成两个loss分别clip:每个loss各允许幅度的更新,叠加后实际更新了2ε→超出信任域

统一后只clip一次:final_adv = outcome + mem → clip一次 → 总更新在ε内

实现极简+计算高效

统一方案:1次前向,1次反向,1次参数更新 → 实现极简+计算高效 → 无需调两个1oss的权重系数

<span>final_adv</span> = outcome_adv + mem_adv  ←   <span>1</span>行加法

独立方案:→2次前向,2次反向,2次参数更新(或需要梯度累积) → 还需要调两个loss的权重系数 →实际上调权重系数 ≈ 调mem_adv的相对幅度(本质相同)

总结:叠加 advantage 本质上等价于带权重的多目标优化,但更稳定、更高效、且天然兼容PPO的clip约束

2.6 梯度

我们来看看整个训练过程中哪些计算有梯度、哪些没有。

✗ 无梯度(detached/frozen):

✗ 无梯度(detached/frozen):
  - Rollout 生成 token                     → SGLang 推理,不保留计算图
  - A1 <span>compute_log_prob</span> <span>(mem_reward)</span>      → detached,仅算数值
  - B2 <span>compute_score</span> <span>(outcome_reward)</span>     → 纯字符串匹配,无张量
  - B4-algo outcome_adv                   → 纯数值运算
  - A2 mem_adv                            → 纯数值运算
  - <span>A3</span> <span>final_adv</span> <span>=</span> outcome + mem          → 常数张量
  - old_log_prob                          → detached 快照
  - ref_log_prob                          → 冻结模型

☑ 有梯度(唯一来源):

☑ 有梯度(唯一来源):
  PPO Update 中:
  new_log_prob = actor<span>.forward</span>(response_ids)    ← 当前 actor 前向
                 ↑ 这是唯一参与计算图的量

  ratio = <span>exp</span>(new_log_prob - old_log_prob)      ← 梯度流经 new_log_prob
  loss = <span>-mean</span>(final_adv × clip(ratio) × <span>mask</span>)  ← 标量
   + KL_coef x <span>f</span>(new_log_prob, ref_log_prob)
   
  loss<span>.backward</span>() 
  ↓ 梯度方向:
  loss → ratio → new_log_prob → actor 参数 (weights, biases, embeddings) 
  optimizer<span>.step</span>() →更新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)衡量两个分布的"距离":

  • = 0:actor和ref完全一样(没学到任何东西)
  • 很大:actor和ref差别巨大(可能过度优化/rewardhacking)

在loss中:

total_loss = policy_loss + KL_coef x KL

policy_loss想让模型<span>"往高reward方向走"</span> → 拉离<span>ref</span> 

KL_penalty想让模型<span>"别离ref太远"</span> → 拉回<span>ref</span>

两者对抗 → 模型在"提升"和"稳定"之间找到平衡

在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):

<span># 普通KL(k1,有偏梯度): </span>
kl ≈ log π_ref - log π_new

<span># Low-variance KL estimator</span>
<span># 代码中用的不是直接的 KL,而是 low-variance 近似:</span>
    <span># low_var_kl(k3,无偏梯度):</span>
    <span>kl</span> = ratio_ref - log(ratio_ref) - <span>1</span>
    where <span>ratio_ref</span> = exp(new_log_prob - ref_log_prob) = π_actor / π_ref

    当 <span>ratio_ref</span> = <span>1</span> (完全一样): kl = <span>1</span> - <span>0</span> - <span>1</span> = <span>0</span> ✅
    当 ratio_ref > 1 (actor概率更高): kl > 0
    当 ratio_ref < 1 (actor概率更低): kl > 0
    → 任何偏离都被惩罚,是一个"弹簧力"把 actor 拉回 ref

为什么用 k3? k3比 k1方差更低,训练更稳定,尤其适合多轮 agent 场景(轨迹长、方差本来就大)。

几个估计器比较如下:

估计器公式性质
k1![]()无偏,方差大,可负
k2![]()有偏,恒正
k3![]() (![]())无偏 + 恒正 + 低方差(Schulman)
low\_var\_kl同k3,再做clamp(-10, 10)数值保护同k3 + 防爆

TransFormer-封面

TransFormer-封面

0xFF 参考

本文使用 markdown.com.cn 排版