资讯动态

长程Agent强化学习实战:从GRPO迁移到FlashREINFORCE的完整指南

发布时间:2026/9/18 3:23:06 来源:尧图企业网站定制
最近在调一个长程 Agent 的强化学习流程用的是 GRPO结果一上真实任务就各种“跑不动”显存先爆掉然后采样效率极低最后 loss 开始乱跳。折腾了一个多星期最后切到 NVIDIA FlashREINFORCE 这套“每个 prompt 只采一条轨迹”的方案才把训练拉回来。这篇文章就把我在这期间踩的坑和最后理顺的思路完整写一遍包括 GRPO 为什么在长程 Agent 上会失效、FlashREINFORCE 的核心机制是什么、以及从工程上怎么迁移和调参。适合正在做 Agent RL、LLM 强化学习训练或者对 GRPO 和 REINFORCE 系算法感兴趣的同学参考前半部分偏原理后半部分全是实操。1. GRPO 为什么能火以及它到底在“省”什么1.1 从 PPO 到 GRPO把 Critic 砍掉的那一步要说清 GRPO 在 Agent 场景为什么跑不动得先搞明白 GRPO 当年是怎么跑起来的。传统强化学习里PPO 是最常用的策略优化方法它需要同时维护一个 Actor策略模型和一个 Critic价值模型。Critic 的作用是估计状态价值用来计算 Advantage也就是“当前动作比平均水平好多少”。Critic 在训练里非常耗资源而且在 LLM 场景下它通常和 Actor 一样大等于你训练一个模型的同时还要再背一个几乎同等规模的模型。DeepSeek 提出 GRPO 时做了一个很关键的动作不要 Critic 了。它改用同一个 prompt 采样出来的多个 response 的奖励做组内比较用组内平均值和标准差把奖励标准化当作 Advantage 近似值。这样就把价值模型的参数量和训练开销整个省掉了。这背后的直觉其实很好理解对于同一个问题如果模型生成了 8 个回答第 1 个回答得分 0.2其他 7 个得分都在 0.8 左右那么第 1 个回答的相对质量就是“明显偏差”。这个相对值就是优势。GRPO 不去学一个“绝对好坏”的价值函数只比较同一批样本的内部差异这在数学推理这类任务上效果很好因为同样的问题采样结果天然可比组内基线足够稳定。1.2 组内相对优势与“不需要反向传播到奖励”GRPO 的计算流程大致是给一个 prompt采样 N 条 response用规则或者奖励模型打分得到一组 reward。然后对每个 response 计算A_i (r_i - mean(r_group)) / std(r_group)这个 A_i 就是组内标准化后的优势。损失函数用类似 PPO 的 clipped surrogate loss外加一个对参考模型的 KL 惩罚L -mean( min(ratio_i * A_i, clip(ratio_i, 1-ε, 1ε) * A_i) ) β * KL(π_θ || π_ref)ratio_i π_θ(response_i | prompt) / π_old(response_i | prompt)这套做法有两个关键的省钱点。第一没有 Critic省掉了一半模型参数量和对应的前向/反向计算。第二奖励不进入计算图也就是说奖励模型不需要可微哪怕你的奖励是调用外部 API 返回一个整数或者是一段代码编译是否通过都没关系。这正好符合 LLM 对齐任务的特点奖励往往来自规则、人工打分、环境反馈根本不可能求导。但注意GRPO 的这套优势估计默认了一个前提同一个 prompt 下面有多条可以横向比较的 response。这个前提在数学题、代码题这类“短上下文单轮生成”任务里非常自然可一旦换到长程 Agent 场景问题就开始冒头了。1.3 GRPO 的舒适区短生成、可并行、奖励可比说实话GRPO 的舒适区非常明确单轮对话、答案长度有限、奖励信号密集且可比较。比如数学竞赛题你采样 8 个解题过程每个过程几千 token最后看答案对错或者过程得分组内标准化完全合理。这种场景下 GRPO 的表现甚至超过 PPO又快又稳。但长程 Agent 是完全不同的物种。Agent 要经历多轮工具调用每轮动作都改变环境状态上下文会随着交互不断变长最终奖励只在任务结束时才出现中间步骤往往没有即时的奖励信号。你让模型写一段代码马上能看到编译是否通过这是短程你让模型去调研一个问题需要搜索网页、阅读文档、整理信息、生成报告整个过程可能上万 token这才叫长程。GRPO 的思想在长程任务上不是“需要修一修”而是根基性的假设被动摇了。2. 长程 Agent 训练GRPO 在哪些地方真正“跑不动”2.1 单条轨迹太长显存峰值直接把人劝退第一个最直观的问题是显存。GRPO 的组内采样意味着同一时刻要处理 N 条完整 response。对数学推理任务每条 response 可能只有 1-2k token8 条也就十几 k显存压力不算大。但长程 Agent 的单条轨迹动辄 8k、16k甚至 32k token 以上因为每一步工具返回的结果都会被追加到上下文里。推理阶段采样还能勉强用高并发 rollout 顶着真正的灾难出现在训练阶段GRPO 更新策略时需要计算每个 token 的 log prob 和梯度。标准的 PyTorch 反向传播需要保存前向计算中的 logits 才能计算交叉熵梯度。假设序列长度 T32k词表大小 V128k单是保存这一层 logits在 bf16 下就是 T×V×2 字节等于 8GB。这还只是一个 token 序列的 logitsGRPO 一次更新要吃 N 条轨迹还要叠加模型中间层的激活值。在 80GB 的卡上稍微做点长程任务OOM 几乎是必然的。我在实际训练中测过8B 模型32k 上下文不采取额外措施时 GRPO 在单卡上根本无法完成一次更新显存直接打满。即使开了 activation checkpointing也只是勉强不崩batch size 被压到极小训练效率惨不忍睹。这还没算上采样后要对齐组内样本长度、padding 带来的额外浪费。2.2 Agent 轨迹只能顺序生成group 采样在工程上近乎不成立GRPO 的第二道坎在 rollout 阶段。数学题可以一次性从模型里采样出多个独立回答因为生成过程不依赖外部环境。但 Agent 不一样模型先输出一个动作环境返回观察结果模型再基于新的上下文输出下一个动作。你根本无法预先把一条完整轨迹的多个“分支”同时生成出来因为后续 token 的分布取决于前一步动作之后的环境反馈。有人会说那我开多个环境并行跑同一个 prompt 的多个副本不就行了理论上可以工程上非常难受。每个环境里的 Agent 轨迹长度不一样一个环境 10 步就结束了另一个环境陷入循环跑出 50 步你要等最慢的那个完成才能凑齐一个 group 做组内标准化。这中间的 GPU 空闲、内存占用、同步开销都是真金白银。更麻烦的是如果连续交互使上下文不停增长同一个 group 内的轨迹可能在“上下文长度”这个维度上严重不一致你会被迫给所有轨迹补齐到相同长度算力浪费非常严重。GRPO 的“组内同步等待”这个特性在长程 Agent 场景里基本就是效率杀手。相比之下单样本策略一个 prompt 对应一条完整轨迹天然不需要等待任何同组样本轨迹一完成立刻就能进入训练队列流水线可以一直跑。2.3 组内相对基准在 Agent 场景下会发生“基准漂移”GRPO 用组内均值和标准差做 Advantage这个做法假设组内样本来自同一个任务、难度一致、采样噪声可控。长程 Agent 任务里这个假设经常不成立。举个例子同一个 prompt8 条轨迹里有一条运气特别好第一步就找到了正确答案所在的网页后面顺风顺水另外 7 条都在无关页面里打转。最后奖励可能是成功的那条 1.0失败的有几条 0.1有几条 0。组内标准化之后0.1 和 0 的差异被放大得很厉害模型会误以为“0.1 的轨迹”是高质量行为但实际上 0.1 只是环境给的同情分中间的全是无效探索。也就是说Agent 轨迹的奖励差异更多来自随机探索路径的分叉而不是策略质量的真实差异。组内相对基准此时不是在“去噪”而是在“制造噪声”。如果你把不同难度的任务打进同一个 batch情况更糟。简单任务的轨迹普遍拿高分难任务的轨迹普遍低分组内标准化把所有分数拉回均值附近模型学到的东西被严重稀释。这是我实际跑 GRPO 时观察到的最明显的现象之一训练 loss 在下降但下游 Agent 评测指标完全不动。2.4 KL 约束与策略偏移的二次伤害还有一个比较隐蔽但影响很大的问题关于 old policy 和 on-policy 假设。GRPO 采样出一批 response 之后用旧策略的 log prob 计算重要性采样比率然后做多轮优化。理论上这是带 ratio 截断的 off-policy 更新容忍度有限。短任务里 sample efficiency 高偏差小长程任务里一条轨迹动辄上万 token策略稍微更新一点整条轨迹的 log prob 就会发生显著变化。你会发现 ratio 很容易就顶到 clip 边界大量样本的梯度被截断实际有效更新非常有限。再加上 KL 惩罚如果你把 KL 系数调低策略快速偏移旧轨迹的 log prob 严重失真重要性采样估计的偏差会被持续放大如果把 KL 系数调高策略几乎又不动Agent 探索能力被锁死。这个平衡点在不同 task 之间差异极大热插拔式的调参几乎无法收敛。很多团队在长程 Agent 上跑 GRPO 跑了半天没效果问题往往不在算法实现而在这个 on-policy 假设被拉得太远。2.5 稀疏奖励让优势估计雪上加霜长程 Agent 的另一个特征是奖励稀疏。整个任务跑完最终结果只有成功/失败两种信号中间步骤没有中间奖励。即便你用过程奖励模型也很难覆盖所有中间决策点。GRPO 的组内相对标准差在这种“奖励只有几个离散取值”的场景下提供的信息量非常有限如果 8 条轨迹里 2 条成功、6 条失败Advantage 无非是两组常数模型学不到“我到底在哪一步走错了”。这个问题的根源是组内标准化只能告诉你“这次比那次好”但无法告诉你“好在哪里”。而长程决策恰恰需要精细的 credit assignment。这也是为什么很多做 Agent RL 的团队最终都会转向带过程奖励或者更细粒度 reward shaping 的方案但如果奖励本身不可得策略梯度算法就只能依赖足够的采样多样性和低方差的梯度估计而这恰恰是 GRPO 在长程场景里最薄弱的地方。3. FlashREINFORCE 的思路把显存问题换成计算问题3.1 REINFORCE 为什么在长程场景反而更合适NVIDIA 提出的 FlashREINFORCE名字里明确写了 “REINFORCE”这是蒙特卡洛策略梯度最朴素的形式。它不依赖 Critic不依赖组内多条样本每个 prompt 只采一条完整轨迹用轨迹的整体回报来更新策略。从理论上看REINFORCE 是真正的 on-policy 无偏估计只要你一直用当前策略采样梯度就是对的。你可能会想这不就是最“原始”的强化学习算法吗方差那么大为什么长程 Agent 反而需要它答案很实在长程 Agent 场景里组的建设成本太高而单条轨迹的信息量又很大。与其强行凑一个组内的相对优势不如老老实实用单条轨迹的绝对回报配合足够大的 batch 和适当的 baseline 来压方差。REINFORCE 虽然方差大但它的每一步都不浪费没有同步等待没有组内 padding轨迹一完成就能进训练。在工程效率和理论正确性上它反而比 GRPO 更匹配长程 Agent。NVIDIA 的 FlashREINFORCE 就是在这个思路上把工程短板补上了。它的核心卖点是让单条超长轨迹的 REINFORCE 训练在显存上可行同时通过重算机制把存储成本转移到计算上。3.2 核心点log prob 不存了backward 时重新算FlashREINFORCE 最关键的工程机制是 rollback也就是反向传播时重算 log prob。传统训练的反向传播需要保存前向计算中的 logits这是显存爆炸的大头。FlashREINFORCE 的做法是前向传播时只保存必要的信息比如 input_ids、position_ids、attention_mask 这些轻量张量而把每一层的激活值、最终的 logits 通通丢掉。等到反向传播的时候额外做一次前向把 logits 重新算出来再接着算梯度。这个思路和 FlashAttention 非常像。FlashAttention 在反向传播时不保存完整的 attention 矩阵而是用保存在 SRAM 里的 softmax 统计量重新算一遍 attention 分数。FlashREINFORCE 把同样的哲学用在了整个模型上用一次额外的前向计算换取巨额的显存节省。最终效果是把存储和序列长度成正比的 logits 矩阵摊薄到近乎 O(1) 的常数空间让 32k 甚至 64k token 的长轨迹在单卡上训练成为可能。代价当然也有训练时间变长了大约多出一次完整的前向开销。但在长程 Agent 任务里时间换显存的交换通常非常划算因为显存是硬约束多一点计算时间完全可以通过并行资源弥补。3.3 和 FlashAttention 的配合显存换计算思路的合并这里有个容易忽略的点FlashREINFORCE 对 FlashAttention 是有依赖的。如果 attention 本身还在用朴素实现前向计算保存的 attention 矩阵又会让显存爆掉那你重算 logits 省出来的空间根本不够用。FlashAttention 已经把 attention 那部分激活值省掉了FlashREINFORCE 再把剩余的大头 logits 省掉两者叠加才能在超长上下文下真正跑起来。我自己的理解是FlashREINFORCE 把整个训练前向链路的显存瓶颈一个接一个拆掉了attention 层的 O(T²) 显存由 FlashAttention 解决输出层和交叉熵的 O(T×V) 显存由 log prob 重算解决中间层激活值由 activation checkpointing 解决。这三板斧下去长序列训练才从“几乎不可能”变成“可行”。这里补充一个工程细节反向传播重算 log prob 时需要确保 dropout mask、layer norm 等随机/统计路径完全一致。实际操作中通常用固定随机种子或者保证模型处于 eval 状态来锁定 dropout 行为否则重算出来的 logits 和你前向时的 logits 对不上梯度就乱了。这是我在迁移过程中踩过的一个隐蔽的坑。3.4 单样本下的 baseline 与方差控制“每个 prompt 只采一条”听起来很美好但 REINFORCE 的方差问题不能回避。单样本情况下没有组内样本可以做标准化你怎么计算 Advantage我在实践中主要用三种方式处理。第一种是在一个 batch 里用多条轨迹的 reward 做减均值除标准差把 Advantage 标准化到合适尺度。这本质上把“组”从同一个 prompt 换成了同一批训练数据规模更大、更稳定实现也简单。第二种是用参考序列的 reward 均值做 baseline比如维护一个滑动平均的 reward 估计Advantage r - running_mean。第三种是引入一个轻量价值网络如果你实在需要更低的方差可以加一个很小的 MLP 头预测 return但这就失去了 GRPO 不要 Critic 的优势一般不建议首选。我在实际任务里用的是 batch 标准化加滑动窗口 baseline 的组合效果比较稳。单条轨迹的绝对 reward 数值可能在不同任务间差异很大但经过 batch 标准化后梯度的尺度基本保持稳定训练不容易爆。FlashREINFORCE 的思路在其实现在我这里表现为一个设计原则你并不需要复杂的优势估计器只要梯度无偏、方差可控、显存能跑长程 Agent 的 RL 就能收敛。4. 实操从 GRPO 迁移到 FlashREINFORCE 的配置清单4.1 前置条件与工具选择如果你想把代码从 GRPO 迁到 FlashREINFORCE 这套思路先确认几个前置条件。第一模型底座要支持 FlashAttention这是长序列训练的基本前提。第二训练框架要支持自定义 backward hook或者直接使用已经实现了 log prob 重算的库比如 NeMo Aligner 里相关的参考实现。第三你的 rollout 框架要支持动态 batch也就是轨迹完成一条送一条而不是等一个 group 齐了再送。工具链上训练侧建议用 Megatron-LM 或者 NeMo Aligner 这种对长序列和模型并行支持好的框架rollout 侧可以用 vLLM 或 SGLang它们对长上下文推理的优化比较成熟。如果你的基础设施比较轻也可以自己写一个简化的训练循环把核心的重算逻辑通过自定义 autograd Function 实现。下面我会给一个最小实现的结构做参考。4.2 训练循环的核心结构FlashREINFORCE 风格的训练循环最关键的是那个自定义的 log prob 计算节点。伪代码大致如下import torch def compute_advantage(trajectory_rewards, baselineNone): # trajectory_rewards: [batch_size] # 返回标准化后的优势 adv trajectory_rewards - baseline adv (adv - adv.mean()) / (adv.std() 1e-8) return adv class RollbackLogProb(torch.autograd.Function): 前向时不保存 logits反向时重算模型得到 logits 从而避免存储 [T, V] 大小的张量。 staticmethod def forward(ctx, input_ids, model, temperature1.0): # 这里只是说明思路实际代码里需要用 model 前向计算得到 logits with torch.no_grad(): logits model(input_ids)[logits] # 长序列时不要保留 log_probs logits.log_softmax(dim-1) # 只保存稀疏的 log prob不保存完整 logits ctx.save_for_backward(input_ids) ctx.model model return log_probs staticmethod def backward(ctx, grad_output): input_ids, ctx.saved_tensors model ctx.model # 反向时重新前向得到 logits 并计算梯度 logits model(input_ids)[logits] log_probs logits.log_softmax(dim-1) # 这里需要手动实现交叉熵类梯度 grad_logits ... return grad_logits, None, None真实实现会复杂很多要处理并行策略、激活重算、梯度累积、混合精度等问题。但核心逻辑就是上面这样前向只开销一次反向时再开销一次省掉的是巨大的 logits 存储。训练主循环可以写成rollout 环境返回一条完整轨迹prompt actions observations。计算整条轨迹的累计奖励和每条 action 的 advantage。用模型对整条序列做前向构造策略梯度 loss -log_prob(action) * advantage。反向传播时触发 rollback 重算得到梯度后更新模型。同时计算对参考模型的 KL作为正则项加入 loss。这里有一个容易忽略的点reference model 的 log prob 不需要梯度所以可以在 rollout 完成后单独跑一遍 reference model 前向只把每个 action token 的 log prob 标量存下来占用的显存是 O(T) 量级完全可接受。KL 惩罚则在整个生成序列上计算不要把 intermediate 状态丢掉否则 KL 估计会有偏差。4.3 超参建议与资源估算迁移之后我跑下来的超参范围大致如下。长程 Agent 任务建议从较小的学习率开始因为单条轨迹包含大量 token一个 step 的信息量远大于短任务。参数建议范围说明learning rate1e-6 ~ 3e-6长序列下梯度信号密集lr 过大会导致策略跳变kl_coef0.01 ~ 0.1长程任务 KL 累积快建议从偏小值开始clip range0.05 ~ 0.2单样本策略下 clip 过大会增加方差过小会拖慢学习grad clip1.0防止长序列梯度爆炸batch size16 ~ 64 条轨迹以轨迹数为单位每条约 8k~32k tokenmax context模型支持的极限长度建议预留至少 10% 的余量防止长度抖动资源估算方面以 8B 模型、上下文 32k、单卡 80GB 为例模型权重约 16GBbf16优化器状态如果用 AdamW 大概是 32GB激活值在 FlashAttention activation checkpointing rollback 重算的加持下控制在 20GB 左右是有可能的。也就是说单张 80GB 卡勉强能跑通但空间很紧实际还是建议用多卡做模型并行或数据并行batch 大一些训练更稳定。70B 级别的模型则需要 8 卡或以上并且要配合张量并行和流水线并行。4.4 数据组织与轨迹 token 利用率还有一个特别容易被忽视的问题数据组织。GRPO 的组内样本天然是“同时同长”的而 FlashREINFORCE 是单条轨迹不停进来需要一套动态 batching 机制把长度相近的轨迹凑到一起减少 padding 浪费。我自己常用的做法是 building batch按轨迹长度排序然后贪心地把长度相近的轨迹放进同一个 batch。这样每条轨迹的 padding 比例可以控制在 10% 以内。另外要特别关注一个指标action token 占比。长程 Agent 轨迹里大量 token 是环境返回的 observation 和中间文本真正参与策略梯度计算的只有模型输出的 action token。如果这个占比太低比如只有 5%那绝大部分计算量都花在了“读上下文”上有效更新信号非常稀疏。这时候要么压缩 observation 长度要么对 observation 部分做截断处理否则训练效率会很差。我一开始没有关注这个指标结果一个 step 算了很多 loss模型就是不学后来统计了 action token 占比才意识到问题。5. 常见问题与排查实录5.1 问题速查表下面这张表是我在迁移和调参过程中实际遇到过的问题以及对应的排查方向。现象可能原因处理方式训练中显存 OOM未开启 log prob 重算activation checkpointing 未打开batch 内 padding 过多确认 rollback 路径生效开启 checkpointing使用动态 batching 控制 paddingloss 爆炸或出现 NaNadvantage 尺度太大lr 过高日志里 KL 值突然飙升对 advantage 做 batch 标准化降低 lr检查 KL 系数是否过小训练不收敛评测指标不动action token 占比过低奖励过于稀疏batch size 太小压缩 observation调整 reward shaping增大 batch 或增大采样并行度rollout 环节 GPU 利用率低轨迹长度差异大静态 batching 等待过长改为动态 batching轨迹完成即送训练不等组反向传播重算时 dropout 不一致重算路径和前向路径的随机行为不同固定随机种子或在 rollout 时把模型切到 eval 模式训练时间比预想慢很多rollback 带来额外前向开销FlashAttention 未生效接受大约 1.5~2 倍的计算开销确认 attention 实现是 flash 版本5.2 几个复盘心得第一个心得是不要在长程 Agent 上一上来就跑完整任务。先用一个 3 到 5 步就能结束的合成 Agent 任务验证整个 RL 流程比如“调用一次工具读取一段信息回答一个问题”。这个任务足够短你可以用普通 GRPO 和 FlashREINFORCE 各跑一遍对比显存占用、训练速度、收敛曲线把 pipeline 的坑先排掉。直接上长任务出了问题你根本分不清是算法问题、奖励问题还是环境问题。第二个心得是日志里一定要把 repeated token 的比例、action token 占比、每条轨迹的 reward 分布单独记录下来。Agent RL 训练里这些指标比 loss 本身更重要。我见过很多训练曲线看起来“在下降”但 Agent 实际行为已经完全退化只是模型学会了用低质量高频动作刷 reward。不记录这些细粒度指标很难定位。第三个心得和 GPU 显存碎片化有关。长序列训练经常出现“显存还有剩余但申请不到大块内存”的情况表现为 nvidia-smi 显示显存占用率很高、但是实际利用率很低。这时候除了检查显存碎片还要看看是不是动态 batching 导致每次 forward 的序列长度波动太大。建议设置torch.cuda.set_per_process_memory_fraction或者调整 PyTorch 的max_split_size_mb有时候能让本来准备 OOM 的训练继续跑下去。5.3 什么时候不要用 FlashREINFORCE最后说一点容易忽视的边界情况。FlashREINFORCE 不是万能的如果你的任务本身就短例如只有几百个 token 的问答、纠错、改写任务GRPO 的组内采样优势依然明显采样成本低、方差小、收敛快没有必要切换到单样本 REINFORCE。如果你的任务允许离线采集大量轨迹并且你可以接受 off-policy 更新那用带重加权的离线 RL 方法可能是更好的选择因为可以反复利用已有轨迹效率更高。FlashREINFORCE 最适合的场景就是那种轨迹超长、上下文持续增长、无法预采样、必须 on-policy 逐步交互的 Agent 任务。只有在这种场景下“每个 prompt 只采一条”才是相较 GRPO 最合理的选择。就我个人这段实操经历来说最大的收获是理解了“算法好不好用取决于它的假设和你的场景是否匹配”。GRPO 的组内相对优势在数学推理上是天才设计但放到长程 Agent 上摇身一变成了显存负担和同步瓶颈。FlashREINFORCE 看起来朴素反而因为抓住了长序列训练的硬约束把问题做成了工程上可解的题。如果你也在长程 Agent 的 RL 训练里挣扎不妨按这套思路先跑通一个小任务再逐步放大应该能少走很多弯路。

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价