Перейти к содержимому
Оригинал
罗西的思考· 罗西的思考·· 13 дней назадОценка ИИ22

MemPO 源码学习笔记(4):Rollout 实现细节

Оригинальный заголовок: [Agent Memory / 强化学习] MemPO源码学习笔记 ---(4)--- Rollout实现细节

Заголовок и краткое изложение на выбранном языке ожидают перевода.

Краткий обзор ИИ

MemPO 源码笔记解析 Rollout 实现细节:mem_sys_prompt_ids 是第一轮 system prompt 加用户问题的深拷贝,作为构建 mem_traj 的干净前缀,使每轮 P_mem 都从原始问题出发评估。

Полный текст

Полный текст на выбранном языке ожидает перевода. Пока показан оригинал.

Illustration accompanying this article

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

  • 0x00 概要

  • 0x01 回顾

  • 0x02 mem_sys_prompt_ids

    • 2.1 mem_sys_prompt_ids 定义

    • 2.2 mem_sys_prompt_ids作用

  • 0x03 ans_mask

    • 3.1 位置

    • 3.2 ans_mask 的精确构造过程

  • 0x04 threshold

    • 4.1 作用

    • 4.2 位置

    • 4.3 过滤的实际意义

    • 4.4 示例

    • 4.5 作用域

    • 4.6 问题

  • 0x05 Misc

    • 5.1 16个样本

    • 5.2 mem_rewards_idx_list

    • 5.3 首尾

  • 0xFF 参考

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 = tokenize("\n<think>...\n</think>\n<answer>\nKathryn Bigelow\n</answer>") ans_mask: [0,0,0,0,1,1,1,1,1,1,1,1,      0,0, 0, 0]
                         ↑Kathryn Bigelow 的token  ↑\n</answer>的4个token

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



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

      具体示例(Round 3)如下:

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

        这就是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([
                  {"role":"system", "content":"You are a helpful assistant..."},
                  {"role":"user","content":"Who directed the 2o1o Best Picture?"}
              ])
          ) # 用于构建mem_traj 时作为"干净前缀"与<mem> 摘要拼接

          token组成如下:

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

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

            2.2 mem_sys_prompt_ids作用

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

              # tool_agent_loop.py
              mem_traj_ids_list.append(agent_data.mem_sys_prompt_ids + response_mem_ids)
              mem_traj=[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 + <mem>摘要</mem>]    → P_mem ← 仅摘要

                关键特点

                  每一轮的 mem_traj 都共享同一个mem_sys_prompt_ids(第一轮的prompt) 
                  → 无论到了第几轮,"上下文起点"永远是"系统提示+原始问题" → 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 [bsz, seq_len]
                    → 区分 prompt token (=0) vs response token (=1)
                    → 在 PPO loss 中:loss = -mean(adv × ratio × response_mask) 
                    → 不涉及 ans_mask 或 threshold


                    Memory Advantage

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

                    • ans_mask:标记answer_ids 中"核心答案token"的位置

                    • threshold:log(0.5),过滤低置信度token

                      full_ans_mask = ans_mask & (full_logp > threshold) 
                      mem_ans_mask = 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="里德学院" #从数据集取第一个答案 


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

                        Step 2: 单独 token 化 core_response_str

                          core_response_str = "里德学院"  # 只有纯答案文本,无 XML 标签
                          core_response_ids = tokenizer("里德学院").input_ids
                          # 假设:[里,德,学,院] = 4 个 token,len = 4

                          Step 3: 计算 ans_mask

                            ans_mask = np.zeros_like(answer_response_ids)  # 全零
                            ans_mask[-1*(len(core_response_ids)+4):-4] = 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 是 1 token
                                </answer> 是 3 tokens(或tokenizer可能分不同方式)

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

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

                                  answer_response_ids:
                                  [\n  <think> \n I...</think> \n <answer> \n  里   德   学   院  \n    </answer>]
                                   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 个


                                    ans_mask[-1*(4+4):-4] =ans_mask[-8:-4] =1
                                    index: 0 1 2 ... 9 10 11 12 13 14 15 16 17
                                    mask:  0 0 0 ... 0 0  0  0  0  1  1   1  1 0 0 0 0 
                                                                   ↑里德学院    ↑ \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 为 ["Kath", "ryn", " Big", "elow"]
                                      log_prob:  [-0.2,  -3.5,  -0.1,  -2.8]
                                      threshold:  -0.693


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

                                      4.2 位置

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

                                      调用位置如下:

                                          文件:verl/experimental/agent_loop/agent_loop.py
                                          函数:AgentLoopManager.generate_sequences() 的后处理段(即 A1)
                                          路径:A4(收集) → [A1] _postprocess → A2(归一化) → A3(叠加)
                                                           ↑ 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%)

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

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

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

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

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

                                              4.4 示例

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

                                              没有threshold:

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


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

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

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

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

                                                    threshold =log(0.5) ≈ -0.693


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


                                                    full_ans_mask = ans_mask AND (full_logp > threshold) 
                                                    mem_ans_mask = ans_mask AND (mem_logp > threshold)

                                                    4.5 作用域

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

                                                      threshold = math.log(0.5)  # -0.693


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


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


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

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

                                                        答案="Reed College"(3 tokens:Re, ed, GCollege)


                                                        情景:full_traj 对"GCollege"很确信,mem_traj不确信
                                                            full_logp:[Re=-0.2,ed=-0.3,GCollege = -0.1] → 全部 KEEP
                                                            mem_logp:[Re=-0.5,ed=-0.4,GCollege = -1.5] → GCollege 被 FILTER 


                                                            P_full基于3个token(token 0,1,2)的均值
                                                            P_mem基于2个token(token 0,1)的均值(跳过了GCollege) 


                                                            P_mem的分母减小(仅2个有效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的多条轨迹来计算组内统计量:
                                                              group_mean = mean([score_1, score_2,...,score_16])
                                                              group_std = std([score_1, score_2,...,score_16])
                                                              adv_i =(score_i -mean) / std


                                                              如果只有1条→无法归一化 
                                                              16条→ 统计量估计相对稳定


                                                          例子:
                                                              Question: "Who directed Inception?"
                                                              轨迹1: search("Inception")→答对→score=1
                                                              轨迹2: search("2010 film")→答错→score=0
                                                              轨迹3: search("Inception director")→ 答对→ score=1
                                                              ......
                                                              轨迹16:search("Nolan movies")→答对 → score=1 


                                                              mean=0.75,std=0.43
                                                              轨迹1:adv=(1-0.75)/0.43=+0.58 (鼓励)
                                                              轨迹2:adv =(0-0.75)/0.43=-1.74(抑制)


                                                          这个值是可配置的,在 run_train.sh 中通过 actor_rollout_ref.rollout.n=16 设置。    

                                                          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。

                                                          0xFF 参考

                                                          Источник: 罗西的思考 · mp.weixin.qq.com