[Agent Memory / 强化学习] MemPO源码学习笔记 ---(4)--- Rollout实现细节

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

适合复现MemPO或做Agent记忆RL训练的工程师,可快速定位P_mem计算的三处关键实现,并提前规避硬编码掩码与阈值带来的坑。

\[Agent Memory / 强化学习\] MemPO源码学习笔记 ---(4)--- Rollout实现细节 ---------------------------------------------------------

0x00 概要

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

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

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

MemPO 的信息如下:

本篇看一些实现的细节,主要是rollout方面的实现细节。

0x01 回顾

我们首先回顾下。

Rollout只生成一种轨迹一一一完整的多轮对话轨迹。full_traj和mem_traj是从这条轨迹中提取/构造出来的。

构建 mem_traj 的方案,即mem_traj 的最终token 组成如下:

4-mem_traj 的最终token 组成

4-mem_traj 的最终token 组成

ans_mask 会对答案 token 做 mask,其含义:只关注"核心答案内容"的 log_prob,忽略/ 标签token。

answer_ids <span>=</span> tokenize(<span>"<span>\n</span><think>...<span>\n</span></think><span>\n</span><answer><span>\n</span>Kathryn Bigelow<span>\n</span></answer>"</span>) ans_mask: [<span>0</span>,<span>0</span>,<span>0</span>,<span>0</span>,<span>1</span>,<span>1</span>,<span>1</span>,<span>1</span>,<span>1</span>,<span>1</span>,<span>1</span>,<span>1</span>,      <span>0</span>,<span>0</span>, <span>0</span>, <span>0</span>]
                     <span>↑</span><span>Kathryn</span> <span>Bigelow</span> 的token  <span>↑</span>\n<span></</span>answer<span>></span>的<span>4</span>个token

threshold 会过滤mask,其含义:忽略模型"完全没把握"的token

<span>full_ans_mask</span> = ans_mask AND (full_logp > log(<span>0.5</span>)) 
<span>mem_ans_mask</span> = ans_mask AND (mem_logp > log(<span>0.5</span>))

具体示例(Round 3)如下:

<span><</span><span>|</span>im_start<span>|</span><span>></span><span>system</span>
 You <span>are</span> a helpful assistant.<span><</span><span>|</span>im_end<span>|</span><span>></span> 
<span><</span><span>|</span>im_start<span>|</span><span>></span><span>user</span>
 乔布斯在哪所大学读书?<span><</span><span>|</span>im_end<span>|</span><span>></span>
<span><</span><span>|</span>im_start<span>|</span><span>></span>assistant 
<span><</span>mem<span>></span>
    乔布斯曾就读于俄勒冈州里德学院,<span>1972</span>年入学, <span>6</span>个月后学,但仍在校旁听书法课。
    最终答案应该是里德学院(Reed College) 
<span><</span><span>/</span>mem<span>></span><span><</span><span>|</span>im_end<span>|</span><span>></span>

这就是P_mem的计算上下文:系统提示 + 原始问题 + 当前轮内容,然后用这个上下文计算模型生成正确答案的概率。

接下来我们分析下具体细节。

0x02 mem_sys_prompt_ids

我们来看看 mem_sys_prompt_ids 包含哪些内容(具体 token 组成)。

2.1 mem_sys_prompt_ids 定义

mem_sys_prompt_ids=第一轮生成前的初始prompt_ids的深拷贝,即system prompt + 用户问题(不含任何多轮对话历史)。具体内容是:

mem_sys_prompt_ids = tokenize(
    apply_chat_template([
        {<span>"role"</span>:<span>"system"</span>, <span>"content"</span>:<span>"You are a helpful assistant..."</span>},
        {<span>"role"</span>:<span>"user"</span>,<span>"content"</span>:<span>"Who directed the 2o1o Best Picture?"</span>}
    ])
) <span># 用于构建mem_traj 时作为"干净前缀"与<mem> 摘要拼接</span>

token组成如下:

<span>|</span> 部分                         <span>|</span> 内容                        <span>|</span>
────────────────────────────────────────────────────────────
<span>|</span><span><</span><span>|</span>im_start<span>|</span><span>></span><span>system</span>\n         <span>|</span> 系统角色起始                  <span>|</span>
<span>|</span> You <span>are</span> a helpful assistant <span>|</span> 系统提示文本                  <span>|</span>
<span>|</span> <span><</span><span>|</span>im_end<span>|</span><span>></span>\n                <span>|</span> 系统段结束                    <span>|</span>
<span>|</span> <span><</span><span>|</span>im_start<span>|</span><span>></span><span>user</span>\n          <span>|</span> 用户起始                     <span>|</span>
<span>|</span> 原始问题文本                  <span>|</span> 训练数据中的question          <span>|</span>
<span>|</span> <span><</span><span>|</span>im_end<span>|</span><span>></span>\n                <span>|</span> 用户段结束                    <span>|</span>
<span>|</span> <span><</span><span>|</span>im_start<span>|</span><span>></span>assistant\n     <span>|</span> 助手起始(generation prompt)  <span>|</span>

注意:prompt_ids 是第一轮开始时的初始 prompt,此时 messages 只包含系统提示 + 用户问题,没有任何历史搜索结果或 内容。

2.2 mem_sys_prompt_ids作用

作用:构建mem_traj时作为"干净的前缀",与摘要拼接:

<span># tool_agent_loop.py</span>
mem_traj_ids_list.append(agent_data.mem_sys_prompt_ids + response_mem_ids)
<span>mem_traj</span>=[system + question]        +    [<mem>摘要内容</mem>]
               ↑ mem_sys_prompt_ids        ↑ response_mem_ids

为什么需要它

在A1中,要对比 full_traj 和 mem_traj,两者必须有相同的"起点"(system+question),才能公平比较。而 mem_sys_prompt_ids 提供这个相同起点。

full_traj= [system +question +多轮完整对话历史]      → P_full ← 完整历史
mem_traj = [system + question + <span><span><<span>mem</span>></span>摘要<span></<span>mem</span>></span></span>]    → P_mem ← 仅摘要

关键特点
每一轮的 mem_traj 都共享同一个<span>mem_sys_prompt_ids</span>(第一轮的prompt) 
→ 无论到了第几轮,<span>"上下文起点"</span>永远是<span>"系统提示+原始问题"</span> → mem_traj不累积历史,每轮都重新从原点开始评估

这使得P_mem真正测量的是:"仅凭这一条,模型能从原始问题出发回答正确吗?"

随着轮次推进,prompt_ids会不断增长(加入工具结果等),但mem_traj 需要的始终是"最初的 system + question" → 必须在第一轮就deepcopy保存。

小结

简言之:mem_sys_prompt_ids 是Memory Reward计算中 "如果模型只看问题+摘要” 这个假设条件的实现。这样可以让模型在"只看问题+摘要" vs "看完整历史"两种条件下预测答案,比较概率差异。

我们接下来介绍 ans_mask 和 threshold。

0x03 ans_mask

3.1 位置

ans_mask 和 threshold 都不作用于 Outcome Advantage。它们仅作用于 Memory Advantage 路径。

Outcome Advantage

Outcome Advantage 路径中的"mask" 作用如下:

只用 response_mask <span>[bsz, seq_len]</span>
→ 区分 prompt token (=<span>0</span>) vs response token (=<span>1</span>)
→ 在 PPO loss 中:loss = <span>-mean</span>(adv × ratio × response_mask) 
→ 不涉及 ans_mask 或 threshold

Memory Advantage

Memory Advantage 路径中的 mask 和 threshold (A1):

  • ans_mask:标记answer_ids 中"核心答案token"的位置
  • threshold:log(0.5),过滤低置信度token
<span>full_ans_mask</span> = ans_mask & (full_logp > threshold) 
<span>mem_ans_mask</span> = ans_mask & (mem_logp > threshold)

→ 用于计算P_full和P_mem → 产出mem_reward

两条路径对比:

  • Outcome: response_str → em_check → {0,1} (无mask / threshold)
  • Memory: log_prob → ans_mask × threshold过滤 → P_mem - P_full

3.2 ans_mask 的精确构造过程

ground_truth = "里德学院" core_response_ids = [里,德,学,院]→ len = 4

Step1:构造完整的"答案序列"
ground_truth_text<span>=</span><span>"里德学院"</span> #从数据集取第一个答案 

answer_response_str <span>=</span> (
    <span>"<span>\n</span><think><span>\n</span>"</span>
    <span>"I have sufficient information to provide the final answers.<span>\n</span>"</span>
    <span>"</think><span>\n</span>"</span>
    <span>"<answer><span>\n</span>"</span>
    <span>"里德学院<span>\n</span>"</span>   # ground_truth_text
    <span>"</answer>"</span>
)

Step 2: 单独 token 化 core_response_str
<span>core_response_str</span> = <span>"里德学院"</span>  <span># 只有纯答案文本,无 XML 标签</span>
<span>core_response_ids</span> = tokenizer(<span>"里德学院"</span>).input_ids
<span># 假设:[里,德,学,院] = 4 个 token,len = 4</span>

Step 3: 计算 ans_mask
<span>ans_mask</span> = np.zeros_like(answer_response_ids)  <span># 全零</span>
ans_mask<span>[-1*(len(core_response_ids)+4):-4]</span> = 1
          ↑                                  ↑
          从倒数第 (core_len + 4) 个位置       到倒数第 4 个位置(不含)

为什么 +4 和 -4?我们看 answer_response_str 的末尾结构:

... \n 里 德 学 院 \n < / answer >
    ↑                ↑
core 开始前面有\n     末尾4个token:n</answer> 这4个不应计入答案

末尾4个token(Qwen tokenizer)对应\n,即['\n','</','answer','>']。实际上\n被token化后恰好是4个token(硬编码假设):

\n 是 <span>1</span> token
</answer> 是 <span>3</span> <span>tokens</span>(或tokenizer可能分不同方式)

这些是格式标签token,不是答案内容本身,排除它们可以确保只评估模型对核心答案内容的预测能力。

因此,得到具体标记结果如下:

answer_response_ids:
[\n  <span><<span>think</span>></span> \n I...<span></<span>think</span>></span> \n <span><<span>answer</span>></span> \n  里   德   学   院  \n    <span></<span>answer</span>></span>]
 0    1..N           N+1..M      M+1     M+2 M+3 M+4  M+5  M+6 M+7    末尾4个  

ans_mask:[0    0...0     0...0    0  0  0  1  1  1  1   0  0  0  0]
                                           ↑ 从-8到-5    ↑末尾4个保持0
                                           (以4个core token为例)

完整示例(具体数字)

ground_truth = "里德学院"

core_response_ids = [里,德,学,院]→ len = 4

answer_response_str token 序列 (假设共18个 token)如下:

位置:0   1   2   3 ... 13   14   15 16 17
     \n <thi nk  \n    \n  <ans wer> \n 里 德 学 院 \n </ ans wer >
                        ↑ 这里开始     ↑ core 4 个   ↑末尾 4 个
                        
<span>ans_mask[-1*(4+4):-4] =ans_mask[-8:-4] =1</span>
<span>index: 0 1 2 ... 9 10 11 12 13 14 15 16 17</span>
<span>mask:  0 0 0 ... 0 0  0  0  0  1  1   1  1 0 0 0 0 </span>
                               ↑里德学院    ↑ \n<answer>
                               这4个是1     这4个是0

关键约束与潜在风险如下:

要素内容
-4 硬编码假设\\n恰好=4个token
适用条件Qwen tokenizer 中 \\n 为 1 token, 为 3 token (可能是</,answer,>)
风险不同 tokenizer 可能 分词结果不同,导致 mask 偏移,即如果 tokenizer 对 n 的分词不是恰好 4个 token(不同 tokenizer、不同语言),答案mask 会错位,奖励计算错误。
正确效果只有纯答案文本(无 XML 标签)的 token 参与概率计算

这样设计的原因:计算 P(答案丨上下文)时,不希望 \n 这些格式 token 干扰概率估算,只关注实际答案词的预测概率。

0x04 threshold

4.1 作用

threshold 在Memory Reward路径(A1)中使用(对full_logp和mem_logp各自独立过滤)。作用是过滤掉模型"完全没信心"的answer token(如人名的中间子词),避免噪声token拉低P_mem和P_full的区分度。

注意:threshold 与 Outcome 路径无关——Outcome 路径(B 系列)的 em_check 是字符串匹配,不涉及任何概率计算或 threshold。

threshold 的特点如下:

方面内容
主要目的过滤掉模型完全不懂的token,避免随机噪声
效果让mem\_reward 聚焦于"有意义的答案token"
副作用两边过滤不同token→P\_mem可能被高估
硬编码风险prob=50%是拍脑袋的阈值,没有消融实验支撑
改进方向可以改为min(full\_logp,mem\_logp) > threshold,确保同一token 才对比

以下面为例,因为 "ryn" 和 "elow" 无论给什么上下文都难预测(子词特性)。如果不过滤,这些噪声 token 会拉低 P_mem 和 P_full,导致 P_mem - P_full ≈ 0(两边都被噪声淹没)。

例如答案 "Kathryn Bigelow",我们得到:

tokenize 为 [<span>"Kath"</span>, <span>"ryn"</span>, <span>" Big"</span>, <span>"elow"</span>]
<span>log_prob:  [-0.2,  -3.5,  -0.1,  -2.8]</span>
<span>threshold:  -0.693</span>

过滤后:   [-0.2,   x,    -0.1,  x   ]  ← 只保留 <span>"Kath"</span> 和 <span>" Big"</span>
                  丢弃          丢弃

4.2 位置

threshold 在 A1_postprocess 中使用,属于 Memory Reward 路径(A 路径)。

调用位置如下:

  文件:verl/experimental/agent_loop/agent_loop<span>.py</span>
  函数:AgentLoopManager<span>.generate_sequences</span>() 的后处理段(即 A1)
  路径:<span>A4</span>(收集) → <span>[A1]</span> _postprocess → <span>A2</span>(归一化) → <span>A3</span>(叠加)
                   ↑ threshold 在这里

调用链如下:

  • ① rollout 完成 → 收集到 full_traj_list, mem_traj_list
  • ② A1: compute_log_prob(2N条) → 得到 full_logp, mem_logp
  • ③ threshold = math.log(0.5) ← 这一步
  • ④ 过滤 + 计算 P_mem - P_full
  • ⑤ 结果存入 mem_rewards → 流向 A2, A3

4.3 过滤的实际意义

情景A:模型"认识"这个答案token(prob>50%)

<span>full_logp</span> = -<span>0.3</span> → KEEP(full_traj 能预测) 
<span>mem_logp</span> = -<span>0.4</span> → KEEP(mem_traj也能预测) 
→ 两边都参与计算,正常对比

情景B:full_traj 能预测但mem_traj不能

<span>full_logp</span> = -<span>0.3</span> → KEEP
<span>mem_logp</span> = -<span>2.0</span> → FILTER(过滤掉)
→ 只有P_full的分子增大,P_mem的分子不增大 → 实际效果:P_mem的均值被计算为"跳过这个token"

情景C:两边都不认识这个token

<span>full_logp</span> = -<span>5.0</span> → FILTER 
<span>mem_logp</span> = -<span>6.0</span> → FILTER
 → 这个token在两边的概率计算中都被排除 → 对比差异 = 0(不干扰信号)

4.4 示例

比如,假设答案 = "亚硫酸盐沉淀反应中间体”(罕见术语)

没有threshold:

<span>full_logp</span>(亚) = -<span>8.0</span> prob = <span>0.0003</span>
<span>mem_logp</span>(亚) = -<span>9.0</span> prob = <span>0.0001</span>
P_full 均值 ≈ <span>exp</span>(-<span>8.0</span>) = <span>0.0003</span> 
P_mem均值 ≈ <span>exp</span>(-<span>9.0</span>) = <span>0.0001</span>
mem_reward = <span>0.0001</span> -<span>0.0003</span> = -<span>0.0002</span>

◄─── 惩罚仅 <span>0.02%</span>,信号极弱且来自无意义的随机猜测差异

有threshold(过滤掉prob<50%的token)

<span>full_logp</span>(亚) = -<span>8.0</span><-<span>0.693</span> → <span>FILTER</span> 
<span>mem_logp</span>(亚) = -<span>9.0</span> < -<span>0.693</span> → <span>FILTER</span>
→ 这条轨迹的答案token 全被过滤,有效 <span>mask</span> 数 = <span>0</span>
→ P_mem = <span>exp</span>(<span>0</span> /(<span>0</span>+<span>1</span>e-<span>8</span>))<span>exp</span>(<span>0</span>) = <span>1.0</span> 
→ P_full = 同上 ≈ <span>1.0</span>
→ mem_reward = <span>1.0</span>-<span>1.0</span> = <span>0</span>(中性,不产生信号)

我们再对threshold = log(0.5)的过滤效果分析。过滤规则如下:

<span>threshold</span> =log(<span>0.5</span>) ≈ -<span>0.693</span>

只有logp > -0.693(即概率>50%)的token才参与计算 

<span>full_ans_mask</span> = ans_mask AND (full_logp > threshold) 
<span>mem_ans_mask</span> = ans_mask AND (mem_logp > threshold)

4.5 作用域

threshold会作用于 full_logp,mem_logp。但是,两者会各自独立过滤。

<span>threshold</span> = math.log(<span>0.5</span>)  <span># -0.693</span>

<span>full_logp_mask_bool</span> = (full_logp > threshold) <span># 过滤full中低概率token</span>
<span>mem_logp_mask_bool</span> = (mem_logp > threshold) <span># 过滤mem中低概率token</span>
<span>full_ans_mask_bool</span> = ans_mask_bool & full_logp_mask_bool <span># 交集,full 的最终 mask </span>
<span>mem_ans_mask_bool</span> = ans_mask_bool & mem_logp_mask_bool <span># 交集,mem 的最终mask</span>

<span># 只对通过过滤的token计算平均log_prob → 再exp</span>
<span>P_full</span> = exp( sum(logp * mask) / sum(mask) ) 
<span>P_mem</span> = exp( sum(logp * mask) / sum(mask) )

<span># 注意:两者的 mask是独立的,可能不同  →  各自用自己"有信心"的token来估算概率</span>

样例如下,这意味着P_full和 P_mem用的是各自的 logp 来过滤,两边可能过滤掉不同的.token一一一一这是设计意图:每个条件下模型对不同 token的置信度可能不同,各自用自己"有信心"的token来估算概率。

答案<span>=</span>"Reed College"(<span>3</span> tokens:Re, ed, GCollege)

情景:full_traj 对"GCollege"很确信,mem_traj不确信
    full_logp:[Re<span>=</span><span>-0.2</span>,ed<span>=</span><span>-0.3</span>,GCollege <span>=</span> <span>-0.1</span>] → 全部 KEEP
    mem_logp:[Re<span>=</span><span>-0.5</span>,ed<span>=</span><span>-0.4</span>,GCollege <span>=</span> <span>-1.5</span>] → GCollege 被 <span>FILTER</span> 
    
    P_full基于<span>3</span>个token(token <span>0</span>,<span>1</span>,<span>2</span>)的均值
    P_mem基于<span>2</span>个token(token <span>0</span>,<span>1</span>)的均值(跳过了GCollege) 

    P_mem的分母减小(仅<span>2</span>个有效token) → P_mem被"拉高"(分母变小),减轻了惩罚

    !这是一个潜在问题:
    当mem_traj 对某些 token 没有把握时,这些 token 被排除
    导致P_mem计算基于"更容易预测的子集",可能虚高

4.6 问题

threshold=log(0.5)是硬编码超参 ,完全没有配置化

  • 对于不同大小的模型(7B vs 70B),合理值差异很大
  • 训练初期模型很弱,大部分 token 被过滤,P_mem 分子为 0→奖励无意义
  • 训练后期模型强了,几乎不过滤→值失效

0x05 Misc

此处介绍其它细节。

5.1 16个样本

16是actor_rollout_ref.rollout.n的配置值一每个question生成16条独立的rollout轨迹,其含义是:同一个question送入LLM16次 → 每次用不同的随机采样(samplingtemperature>0)→ 得到16条内容不同的多轮对话轨迹

16是GRPO的group size(actor_rollout_ref.rollout.n=16)。GRPO用组内均值和标准差归一化advantage,组太小(如2条)→ 均值/方差估计不准,信号噪声大;组太大(如64条)→ 计算开销大,rollout时间长。

为什么MemPO每个question要生成16条rollout轨迹?

  • 16是常见的平衡点。同时,MemoryAdvantage也受益于大组:每个question约有16x3= 48个。
  • mem_reward值用于归一化,统计更稳定。
GRPO需要同一个question的多条轨迹来计算组内统计量:
    <span>group_mean</span> = mean([score_1, score_2,...,score_16])
    <span>group_std</span> = std([score_1, score_2,...,score_16])
    <span>adv_i</span> =(score_i -mean) / std
    
    如果只有1条→无法归一化 
    16条→ 统计量估计相对稳定
    
例子:
    Question: "Who directed Inception?"
    轨迹1: search("Inception")→答对→<span>score</span>=<span>1</span>
    轨迹2: search("2010 film")→答错→<span>score</span>=<span>0</span>
    轨迹3: search("Inception director")→ 答对→ <span>score</span>=<span>1</span>
    ......
    轨迹16:search("Nolan movies")→答对 → <span>score</span>=<span>1</span> 
                
    <span>mean</span>=<span>0.75</span>,std=<span>0.43</span>
    轨迹1:<span>adv</span>=(<span>1</span>-<span>0.75</span>)/<span>0.43</span>=+<span>0.58</span> (鼓励)
    轨迹2:<span>adv</span> =(<span>0</span>-<span>0.75</span>)/<span>0.43</span>=-<span>1.74</span>(抑制)
    
这个值是可配置的,在 run_train.sh 中通过 <span>actor_rollout_ref.rollout.n</span>=<span>16</span> 设置。    

5.2 mem_rewards_idx_list

mem_rewards_idx_list 中0、1、2分别代表什么?

  • 0=无关token(不在区间内)1=开始位置-2=结束位置

5.3 首尾

第1轮(Round1)为什么不收集mem数据?

  • 第1轮是模型第一次生成,没有之前的多轮历史需要总结,因此不会(也不应该)产生摘要。此时mem_rewards_idx_list 全部填 0。

如果rollout被截断,最后一个没有,系统如何处理?

  • 丢弃该轮的mem 数据。检测方式:start_idxs比end_idxs多一个→删除最后一个start。

TransFormer-封面

TransFormer-封面

0xFF 参考

本文使用 markdown.com.cn 排版