文章摘要
针对长上下文智能体记忆管理难题,本文解析自主记忆策略优化框架MemPO所用的GRPO算法。GRPO是PPO的轻量化变体,无需额外Critic网络,改用组内统计量计算优势。文章对比GRPO与标准PPO的差异,介绍MemPO对GRPO缺陷的改进、选择GRPO的原因与具体实现细节,公开了相关项目资源。

一、项目背景与核心设计思路

当前长上下文智能体的记忆管理一直是行业难点,传统方案往往无法对记忆内容进行有效引导,导致生成的记忆信息冗余或无关,难以真正辅助任务完成。

MemPO作为一种自主记忆策略优化框架,让智能体在每轮交互的开头自主生成记忆内容,这种设计将记忆本身转化为可训练的策略变量,通过强化学习信号端到端地教会模型“该记住什么、如何生成高质量记忆”,完全不需要额外的独立记忆模块。

本次我们将重点解析MemPO中GRPO算法的具体应用,先附上项目基础信息:

  • 论文标题: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算法的核心原理与PPO对比

2.1 GRPO与标准PPO的核心差异

GRPO是PPO算法的一个轻量化变体,其最核心的改进在于优势函数的计算方式:用组内统计量替代了原本需要单独训练的Critic网络。简单来说,GRPO将原本需要学习的价值估计组件,替换为同批次轨迹的均值和标准差计算,这也意味着我们需要为每个问题生成多条采样轨迹才能实现该计算。

标准PPO的优势计算方式为:advantage = V_critic(s) - R(s),需要单独训练一个价值网络来估计每个状态的未来回报,这在长序列任务中会面临诸多挑战。

而GRPO的优势计算则简化为:advantage = (reward - group_mean) / group_std,其中group指的是同一个问题下的多条采样轨迹,通常设置为16条。这种方式完全不需要额外训练Critic网络,省去了额外的参数和训练步骤。

我们可以通过下表直观对比两者的差异:

对比维度 标准PPO GRPO
单问题轨迹数 通常1条 16条(组大小)
Critic网络需求 需要,约7B参数规模 无需额外网络
优势来源 GAE:reward-V(s) (score-mean)/std
额外训练步骤 Critic损失训练 无额外训练
显存占用 actor+ref+critic actor+ref
优势精度 token级,但存在估计误差 轨迹级,无估计误差
适配场景 稠密奖励场景 稀疏/结果型奖励场景

除此之外,GRPO完整保留了PPO的其他核心机制:裁剪的代理损失、重要性采样比例、参考模型KL惩罚、多epoch小批量更新,只是将价值基线替换为了组内相对排名。

2.2 现有GRPO的局限性与MemPO的改进

基础的GRPO算法基于最终答案的正确性计算奖励,并且使用轨迹级别的统一优势值,也就是同一条轨迹内的所有token共享同一个奖励信号。这种设计会导致记忆生成的奖励信号非常稀疏,无法精准指导每一步的记忆生成质量——毕竟最终答案的正确性,无法直接反映交互过程中每一次记忆生成操作的优劣。

针对这个问题,MemPO提出了创新性的优势计算方案:在原本的轨迹级优势之外,额外对每一步生成的记忆内容的信息有效度进行评估,计算得到额外的记忆优势值,从而确保生成的记忆在保持简洁的同时,能够保留对任务最关键的信息。

具体来说,MemPO会同时计算两种优势: 1. 结果优势(outcome_adv):沿用标准GRPO的组内归一化计算方式 2. 记忆优势(mem_adv):针对记忆片段单独设计的归一化优势,仅作用于记忆内容区间 最终的总优势为两者的叠加:final_adv = outcome_adv + mem_adv

三、MemPO选择GRPO而非PPO+Critic的原因

在MemPO的场景中,选择GRPO而非传统PPO+Critic方案有四个核心原因:

3.1 长序列下Critic训练难度极大

标准PPO的价值网络需要为序列中的每个token位置预测未来的累计回报,也就是预测“该状态下后续能否正确完成任务”。但MemPO的序列结构通常是多轮交互的拼接,长度可达数千token,而奖励信号仅在序列的最末尾才会给出,属于典型的稀疏奖励场景。

这种情况下Critic网络会面临三重挑战: - 序列极长,需要超大容量的价值网络才能覆盖完整的上下文 - 奖励极端稀疏,导致价值估计几乎处处为0,无法学到有意义的信号 - 多轮工具交互带来的状态空间复杂度极高,价值估计的噪声非常大

3.2 GRPO适配稀疏的结果型奖励

GRPO采用同组轨迹的均值作为价值基线,非常适配轨迹级别的离散奖励场景。MemPO的奖励恰好是轨迹级别的二元评分,完美匹配GRPO的假设前提。同时GRPO不需要学习价值函数,完全通过组内统计量计算优势,避免了长序列下的价值估计误差问题。

3.3 显著节约计算资源

如果采用PPO+Critic方案,需要额外部署一个与actor规模相同的Critic网络,会带来三重资源开销: - 7B参数的额外网络负载 - Critic网络的前向和反向传播计算 - 显存占用翻倍:actor+ref+critic总计21B参数

而GRPO仅需要actor和ref两个模型,总计14B参数,节省下来的资源可以支持更多的并发轨迹采样,进一步提升训练稳定性。

3.4 记忆奖励的天然基线特性

MemPO中的记忆奖励本身就自带基线:mem_reward = P_mem - P_full,也就是通过对比带记忆和不带记忆的场景下的模型表现得到的信号,不需要额外的Critic来估计价值。如果强行为记忆区间单独训练价值头,会面临严重的长期依赖问题——记忆的好坏取决于整个任务的最终结果,Critic几乎无法准确完成这种长周期的价值估计。

而GRPO的方案可以直接对记忆奖励进行跨轨迹归一化,简单高效地完成优势计算。

四、MemPO中GRPO的具体实现细节

4.1 算法流程拆分

MemPO中的GRPO算法可以分为两个核心环节:

环节一:GRPO优势计算(无梯度环节)

# 结果优势计算
outcome_adv = (score - mean) / std
# 记忆优势计算
mem_adv = (r_t - mean) / std
# 总优势叠加
final_adv = outcome_adv + mem_adv

所有的优势计算都是纯数值运算,不会参与反向传播,属于脱离计算图的常数张量。

环节二:PPO参数更新(带梯度环节)

for epoch in ppo_epochs:
    for mini_batch in shuffle(batch):
        new_log_prob = actor.forward(mini_batch)
        ratio = exp(new_log_prob - old_log_prob)
        loss = -mean(final_adv * clip(ratio)) + KL_penalty
        loss.backward()
        optimizer.step()

可以看到,GRPO只负责确定每个token应该被鼓励还是抑制,以及具体的强度,而具体的模型参数优化仍然由PPO的标准流程完成,包括梯度计算、反向传播和参数更新。

4.2 模型组件说明

MemPO的训练涉及三个核心模型组件:

  • actor(策略模型):正在被训练的大语言模型,比如Qwen2.5-7B,每个PPO迭代步骤都会更新其参数,既用于轨迹采样,也用于PPO更新时的前向计算,初始化来自SFT微调后的模型。
  • ref_model(参考模型):与actor结构完全相同的大模型,但参数完全冻结不更新,初始化为训练开始时的actor快照,也就是SFT模型本身,用于计算KL散度惩罚,防止actor偏离初始策略过远。
  • rollout采样器:使用actor的权重进行自回归解码生成轨迹,通常通过SGLang服务实现高效推理。

三者在损失计算中的角色分别是: - ratio:当前actor与历史actor的概率比值,用于重要性采样 - KL_penalty:当前actor与参考模型的KL散度,用于限制更新幅度 - 总损失:结合优势值的裁剪代理损失与KL惩罚

4.3 优势函数的具体设计

MemPO中的两种优势函数都采用了GRPO风格的归一化方式,但在具体实现上各有侧重:

结果优势(Outcome Advantage)

完全沿用标准GRPO的计算流程: 1. 将同一个问题下的16条轨迹分为一组 2. 计算组内奖励的均值和标准差 3. 对每条轨迹的奖励进行组内归一化得到结果优势

记忆优势(Memory Advantage)

这是MemPO的创新性改进,借鉴了GRPO的组内归一化思想但扩展了归一化范围: 1. 分组范围不仅包括同问题的所有轨迹,还包括所有轮次的记忆片段,总计约48个值 2. 对记忆奖励进行跨轨迹、跨轮次的池化归一化 3. 仅在记忆内容区间内生效该优势值

最终总优势为两者的叠加,共同送入标准的PPO裁剪损失中进行训练。

4.4 前向传播流程

MemPO相比基础GRPO多了一次额外的前向传播,但该次传播会同时批量处理完整轨迹和带记忆的轨迹,实际吞吐开销约为标准old_log_prob计算的1.5~2倍,也是MemPO最主要的训练额外成本。完整的训练前向流程如下:

  • 轨迹生成阶段:通过SGLang进行自回归解码,生成16条采样轨迹,利用KV缓存优化,仅计一次前向开销
  • MemPO专属计算:同时计算完整轨迹和带记忆轨迹的对数概率,一次调用处理所有序列
  • 旧对数概率计算:计算当前actor在轨迹上的旧对数概率,用于PPO的重要性采样比例计算
  • 参考模型对数概率计算:通过ref_model计算轨迹的对数概率,用于KL惩罚项
  • 模型更新阶段:多epoch反向传播更新actor参数,默认ppo_epochs=1

与其他同类型算法相比,MemPO多了一次双轨迹的对数概率计算,总前向传播次数有所增加。

4.5 损失函数设计

MemPO的损失函数完全沿用PPO的标准裁剪代理损失,只是将优势函数替换为了叠加后的总优势:

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

这里需要注意几个关键点: 1. 不是两个独立的损失函数,而是一次前向传播、一次损失计算、一次反向传播完成训练 2. 不同类型的token会接收到不同的优势值:记忆区间的token会同时收到结果优势和记忆优势,而其他区间的token仅收到结果优势 3. 通过response_mask屏蔽掉prompt token,仅对响应token计算损失

损失计算的细节拆解

整个损失计算的输入包括: - new_log_prob:当前actor前向得到的每个token的对数概率 - old_log_prob:轨迹采样时的旧模型对数概率,脱离计算图 - ref_log_prob:参考模型得到的对数概率,冻结模型输出 - final_adv:叠加后的总优势张量 - response_mask:标记哪些是响应token,哪些是prompt token

具体的计算流程可以分为: 1. 计算重要性采样比例ratio = exp(new_log_prob - old_log_prob) 2. 计算裁剪后的代理损失surr = min(ratio × final_adv, clip(ratio,1-ε,1+ε) × final_adv) 3. 对所有响应token取平均得到策略损失 4. 计算KL惩罚项,采用低方差的估计量 5. 总损失为策略损失加上加权后的KL损失

4.6 梯度传递机制

在整个训练流程中,只有部分计算会参与梯度传递: 无梯度的环节: - 轨迹生成的推理,不保留计算图 - 记忆奖励和结果奖励的计算,纯数值运算 - 优势函数的所有计算,均为脱离计算图的常数 - 旧对数概率和参考模型对数概率,均为冻结或脱离计算图的结果

唯一的梯度来源: 仅在PPO更新阶段,actor的前向传播计算new_log_prob时会构建计算图,后续的损失计算和反向传播都会将梯度传递到actor的参数中。也就是说,整个MemPO的梯度仅来自于new_log_prob对actor参数的偏导,所有的优势值、奖励值都只是决定梯度的方向和大小,不会直接参与梯度计算。

4.7 KL约束的作用与实现

KL约束的核心目的是防止模型为了刷取奖励而“跑偏”,确保更新后的策略不会偏离初始策略过远,就像学生在学习时可以改进解题方法,但不能完全脱离原本的知识基础。

从数学角度来看,KL(π_actor || π_ref)衡量了两个策略分布之间的距离:当两者完全相同时KL值为0,当差异过大时KL值会显著上升。在损失函数中,KL惩罚项会与策略损失形成对抗:策略损失希望模型向高奖励方向改进,也就是拉离参考模型,而KL惩罚项希望模型不要偏离参考模型过远,也就是拉回参考模型,最终让模型在“提升”和“稳定”之间找到平衡。

如果没有KL惩罚项,可能会出现一系列问题: - 模型生成固定模板的高奖励记忆,但实际没有有效信息 - 模型对所有问题都使用相同的搜索query,依赖巧合获得高分 - 输出多样性崩溃,所有采样轨迹趋于一致 - 训练过程出现不稳定震荡

MemPO中采用了低方差的KL估计量,这种估计器相比传统估计器方差更低,特别适合长轨迹的多轮智能体场景,并且加入了数值裁剪来防止异常值。

五、为什么选择叠加优势而非双独立损失

有些读者可能会好奇,为什么不分别为结果奖励和记忆奖励设置两个独立的损失函数进行训练?其实这种方案存在诸多问题,而叠加优势的方式更加合理:

5.1 避免梯度冲突

如果采用两个独立损失,当出现“轨迹结果错误但记忆质量优秀”的场景时,两个损失会产生方向相反的梯度:结果损失希望降低记忆token的概率,而记忆损失希望提升记忆token的概率,这种梯度冲突会导致训练不稳定、模型震荡。而叠加优势的方式可以将两者的影响合并为一个明确的梯度方向,得到清晰的优化方向。

5.2 适配PPO的裁剪机制

PPO的裁剪约束是为了限制单步更新的幅度,如果拆分为两个独立损失分别进行裁剪,每个损失的更新幅度都会被限制在ε范围内,叠加后实际的更新幅度会达到2ε,超出信任域范围。而叠加优势后只需要进行一次裁剪,确保总更新幅度始终在ε范围内。

5.3 实现极简且高效

叠加优势的方式仅需要一次前向传播、一次反向传播和一次参数更新,实现逻辑非常简洁,不需要调整两个损失的权重系数。而独立损失方案需要两次前向、两次反向和两次参数更新,还需要额外调整权重系数来平衡两个损失的影响,实现复杂度大大提升。

六、参考资料

  • 原论文:MemPO: Self-Memory Policy Optimization for Long-Horizon Agents
  • 开源代码:https://github.com/TheNewBeeKing/MemPO
  • 模型与数据集:https://huggingface.co/collections/NewBeeKing/mempo
以上内容不代表本平台立场,仅供读者参考