文章摘要
当前基于强化学习的内存管理普遍存在缺乏记忆更新定向引导、记忆质量参差不齐的问题,针对该痛点,本文介绍了创新框架MemPO,它可让模型自主管理记忆,提升记忆实用性。本文深度解析MemPO中rollout模块的实现细节,逐一拆解各关键组件的设计逻辑与代码实现,同时附上相关论文、代码等资源链接。

当前基于强化学习的内存管理方案普遍存在一个核心痛点:缺乏对记忆内容更新的定向引导机制,导致生成的记忆质量参差不齐,难以真正服务于任务目标。MemPO,全称Self-Memory Policy Optimization,正是针对这一问题提出的创新框架,它让模型自主管理自身的记忆存储,并引入了基于信息有效度的记忆层级优势估计方法,引导模型优先保留对任务解决最有价值的记忆内容,从而显著提升记忆的实用性与有效性。

MemPO的核心设计思路非常巧妙:让模型在每一轮对话的开头写入记忆内容,形式上类似“自我对话的草稿纸”,既作为任务记忆的载体,也作为思考过程的一部分。这种设计将记忆模块转化为可训练的策略变量,通过强化学习信号端到端地教会模型“应该记住什么、如何构建记忆”,整个过程无需额外的独立记忆模块,完全整合在模型的原生训练流程中。

本次分享的MemPO相关资源包括:

本文将聚焦MemPO中rollout模块的具体实现细节,逐一拆解各个关键组件的设计逻辑与代码实现方式。

核心概念回顾

在深入讲解具体实现细节前,我们先回顾一下MemPO中rollout的核心逻辑:rollout过程只会生成完整的多轮对话轨迹,后续的full_traj和mem_traj都是从这条基础轨迹中提取和构造而来的。

我们以mem_traj的构成为例,它的token序列结构遵循固定的模式:以系统提示和原始用户问题作为基础上下文,随后拼接上模型生成的记忆摘要内容,最终形成完整的记忆轨迹上下文。

ans_mask的核心作用是对答案相关的token进行掩码筛选,它的设计目标是只保留“核心答案内容”对应的log_prob计算,而忽略思考标签、闭合标签等格式性token的干扰。举个简单的例子,当我们对包含思考和答案的序列进行token化后,ans_mask会将对应正式答案文本的位置标记为1,其余格式token的位置标记为0。

threshold的作用是对掩码后的token进行二次过滤,核心目的是排除模型完全没有把握生成的token,避免低置信度的噪声token影响记忆奖励的计算准确性。具体来说,threshold会过滤掉log概率低于log(0.5)(即模型预测置信度低于50%)的token,确保只保留模型有一定把握的答案token参与后续的概率计算。

mem_sys_prompt_ids的设计与实现

2.1 核心定义

mem_sys_prompt_ids是第一轮生成前的初始prompt_ids的深拷贝,其内容仅包含系统提示词与原始用户问题,不包含任何多轮对话历史。具体来说,它是通过对格式化后的对话消息进行token化得到的,对话消息仅包含系统角色和用户角色的初始提问,示例如下:

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组成结构清晰分为几个部分:系统角色起始标记、系统提示文本、系统段结束标记、用户角色起始标记、原始问题文本、用户段结束标记、助手角色生成起始标记。

需要特别注意的是,这里的prompt_ids是第一轮对话开始时的初始上下文,此时的消息列表中仅包含系统提示和用户的初始问题,没有任何过往的搜索结果或生成的记忆内容。

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

为什么需要这样的前缀?因为在MemPO的A1对比逻辑中,full_traj和mem_traj必须拥有完全相同的上下文起点,才能保证两者的概率对比是公平有效的。full_traj的完整上下文是“系统提示+原始问题+完整多轮对话历史”,而mem_traj的上下文则是“系统提示+原始问题+记忆摘要”,两者的初始上下文必须一致,才能准确对比模型在不同上下文下的答案生成概率。

mem_sys_prompt_ids还有一个关键特点:每一轮的mem_traj都会共享同一个初始的mem_sys_prompt_ids,无论对话进行到第几轮,mem_traj的上下文起点永远是“系统提示+原始问题”,不会累积过往的对话历史。这意味着P_mem的计算本质上是在验证:“仅依靠当前的记忆摘要,模型能否从原始问题出发正确生成答案?”

随着对话轮次的推进,普通的prompt_ids会不断累积新的对话历史和工具返回结果,但mem_traj始终需要以最初的系统提示和原始问题作为起点,因此必须在第一轮对话开始时就完成深拷贝并保存。

简单来说,mem_sys_prompt_ids实现了MemPO中“如果模型仅通过问题和记忆摘要来生成答案”的假设条件,让模型可以在“仅看问题+摘要”和“看完整对话历史”两种上下文条件下分别预测答案,并对比两者的概率差异,从而计算出记忆的有效程度。

ans_mask的精确构造逻辑

3.1 应用范围与基础逻辑

ans_mask和threshold都不会作用于Outcome Advantage路径,仅用于Memory Advantage的计算路径中。我们可以将两种优势计算路径进行简单对比:

  • Outcome Advantage路径:直接通过response_str进行精确匹配得到0或1的结果,不涉及任何掩码或阈值过滤

  • Memory Advantage路径:通过计算模型的log_prob概率,结合ans_mask和threshold进行过滤后,计算P_mem与P_full的差值得到记忆奖励

Outcome路径中仅使用response_mask来区分提示token和回复token,在PPO损失计算中仅对回复token的优势值进行加权,不会涉及ans_mask或threshold。

3.2 完整构造步骤

我们以ground_truth为“里德学院”为例,详细讲解ans_mask的构造步骤:

步骤1:构造完整的答案序列

首先需要构建包含格式标签的完整答案字符串,包括思考标签、闭合思考标签、答案标签、正式答案文本以及闭合答案标签,示例如下:

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>"
)

步骤2:单独token化核心答案文本

将纯答案文本(不包含任何格式标签)单独进行token化,得到核心答案的token序列,例如“里德学院”会被token化为[里,德,学,院],长度为4。

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

步骤3:生成ans_mask数组

首先初始化一个与完整答案序列长度一致的全零数组,然后将对应核心答案文本的位置设置为1,具体的切片规则为:从倒数第(core_len +4)个位置到倒数第4个位置(不包含倒数第4个位置)。这里的+4是为了排除末尾的格式标签token。

ans_mask = np.zeros_like(answer_response_ids) # 全零初始化
ans_mask[-1*(len(core_response_ids)+4):-4] = 1
                ↑ 从倒数第(core_len +4)个位置          ↑ 到倒数第4个位置(不含)

为什么是+4和-4?因为完整答案序列的末尾结构为“\n 里德学院 \n </answer>”,核心答案文本前面有一个换行符,末尾的四个token分别是换行符和闭合答案标签的各个部分,这些都属于格式标签,不应计入核心答案的掩码范围。

以Qwen tokenizer为例,换行符\n会被token化为1个token,而</answer>会被token化为3个token,总共4个格式标签token,因此需要通过-4来排除这些末尾的格式token。

最终的ans_mask只会让核心答案文本的token参与概率计算,排除所有的格式标签token,确保模型的概率计算仅针对正式的答案内容,不受格式标签的干扰。

这里需要注意的是,这种构造方式依赖于特定tokenizer的分词结果,如果更换其他tokenizer,可能会导致末尾格式标签的token数量发生变化,从而导致ans_mask的位置偏移,影响奖励计算的准确性。

threshold的作用与实现细节

4.1 核心作用与设计初衷

threshold仅在Memory Reward路径(A1路径)中使用,用于过滤掉模型完全没有信心生成的答案token,避免低置信度的噪声token拉低P_mem和P_full的区分度。例如,对于人名的子词分词结果,模型可能无法准确预测,这些token就会被threshold过滤掉。

需要特别说明的是,threshold与Outcome路径完全无关,Outcome路径通过字符串精确匹配来判断答案是否正确,不涉及任何概率计算或阈值过滤。

threshold的核心特点包括:

  • 主要目的:过滤模型完全无法理解的token,避免随机噪声干扰奖励计算

  • 实际效果:让mem_reward的计算仅聚焦于模型有一定把握的有效答案token

  • 潜在副作用:由于full_logp和mem_logp会各自独立进行过滤,可能导致两边保留的token不一致,从而让P_mem的计算结果被高估

  • 硬编码风险:当前threshold固定为log(0.5),即50%的置信度阈值,没有经过消融实验验证,不同规模的模型或不同训练阶段的合理阈值可能差异很大

  • 改进方向:可以将阈值修改为min(full_logp, mem_logp) > threshold,确保只有在两个上下文下模型都有一定把握的token才参与对比计算

举个具体的例子,答案“Kathryn Bigelow”会被token化为["Kath", "ryn", " Big", "elow"],对应的log_prob分别为[-0.2, -3.5, -0.1, -2.8],threshold为-0.693,那么过滤后只会保留“Kath”和“ Big”这两个token,其余两个低置信度的token会被排除。

4.2 调用位置与执行流程

threshold在A1_postprocess函数中使用,属于Memory Reward路径的后处理阶段。具体的调用流程如下:

  • rollout过程完成,收集到full_traj_list和mem_traj_list

  • 通过compute_log_prob计算full_logp和mem_logp

  • 设置threshold为math.log(0.5),即-0.693

  • 对full_logp和mem_logp分别进行过滤,结合ans_mask得到最终的有效掩码

  • 计算P_mem和P_full的差值,得到mem_rewards并传入后续的归一化和叠加步骤

具体的代码实现示例如下:

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

4.3 不同场景下的过滤效果

我们可以通过三种不同的场景来分析threshold的过滤效果:

  • 场景A:模型对token有充分把握:full_logp和mem_logp都大于threshold,两个上下文下的token都被保留,正常进行概率对比

  • 场景B:full_traj能预测但mem_traj不能:full_logp大于threshold,mem_logp小于threshold,仅保留full_traj的token计算,此时P_mem的计算会跳过该token,导致P_mem的均值降低

  • 场景C:两边都没有把握:full_logp和mem_logp都小于threshold,该token在两边的计算中都被排除,不会对对比结果产生任何干扰

我们以罕见术语“亚硫酸盐沉淀反应中间体”为例,当没有使用threshold时,P_full的均值约为exp(-8.0)=0.0003,P_mem的均值约为exp(-9.0)=0.0001,mem_reward为-0.0002,信号非常微弱,几乎无法起到有效的引导作用。而当使用threshold过滤后,该token的log_prob远小于-0.693,会被直接过滤掉,此时该轮的所有答案token都被排除,mem_reward为0,不会产生任何信号。

4.4 独立过滤的作用域与潜在问题

threshold会分别对full_logp和mem_logp进行独立过滤,也就是说,full_traj和mem_traj会保留各自置信度高于50%的token,两者的有效掩码可能并不相同。例如,full_traj对“GCollege”这个token有充分把握,但mem_traj对该token的置信度较低,那么full_ans_mask会保留该token,而mem_ans_mask会过滤掉该token。

这种设计的初衷是让每个上下文下的模型都仅使用自己有把握的token来计算概率,但也带来了一个潜在问题:当mem_traj对某些token没有把握时,这些token会被排除,导致P_mem的计算仅基于更容易预测的token子集,从而可能高估P_mem的结果。

4.5 硬编码阈值的问题

当前的threshold=log(0.5)是一个硬编码的超参数,完全没有提供配置化的支持,这会带来几个明显的问题:

  • 不同规模的模型(如7B和70B)的合理阈值差异很大

  • 训练初期模型性能较弱,大部分token都会被过滤掉,导致P_mem的计算分子为0,奖励计算失去意义

  • 训练后期模型性能较强,几乎不会有token被过滤,导致阈值的过滤作用失效

其他实现细节补充

5.1 16个样本的配置逻辑

16是actor_rollout_ref.rollout.n的默认配置值,其含义是每个问题会生成16条独立的rollout轨迹。具体来说,就是将同一个用户问题送入模型16次,每次使用不同的随机采样温度,从而得到16条内容不同的多轮对话轨迹。

这个16也是GRPO算法中的group size,GRPO会使用组内的均值和标准差对优势值进行归一化。如果组内样本数量太少(如2条),则均值和方差的估计会非常不准确,导致信号噪声过大;如果组内样本数量太多(如64条),则会带来较大的计算开销,延长rollout的执行时间。

选择16作为group size是一个常见的平衡点,既能保证统计量的估计稳定性,又不会带来过大的计算开销。同时,Memory Advantage的计算也受益于较大的组规模,每个问题的16条轨迹可以让mem_reward的统计更加稳定。

GRPO的归一化逻辑如下:

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条轨迹,则无法进行归一化计算,因此必须使用多条轨迹。举个具体的例子,假设问题是“Who directed Inception?”,16条轨迹中有12条答对,4条答错,那么group_mean为0.75,group_std为0.43,每条轨迹的优势值会根据这个均值和标准差进行归一化。

这个配置值是可以修改的,通常在run_train.sh脚本中通过actor_rollout_ref.rollout.n=16来设置。

5.2 mem_rewards_idx_list的含义

mem_rewards_idx_list中的三个取值分别代表不同的token类型:

  • 0:无关token,不在记忆区间内的token

  • 1:记忆内容的起始位置token

  • 2:记忆内容的结束位置token

5.3 首尾轮次的特殊处理

为什么第一轮对话不会收集mem数据?因为第一轮是模型第一次生成回复,没有过往的多轮对话历史需要总结,因此不会产生记忆摘要,此时mem_rewards_idx_list中的所有值都为0。

如果rollout过程被提前截断,导致最后一个记忆片段没有闭合标签,系统会如何处理?此时会直接丢弃该轮的记忆数据,检测方式为:如果start_idxs的数量比end_idxs多一个,则说明最后一个记忆片段没有闭合,需要删除最后一个start标记。

以上内容不代表本平台立场,仅供读者参考