GiGPO 到底靠什么工作:Agent 训练里的 credit assignment 廉价补丁

GiGPO 把 GRPO 从"一整条轨迹一个 advantage"升级为"episode 组 + step 组 相加",在 ALFWorld/WebShop 上比 GRPO 高 12/9 个百分点。但它的核心机制既不是嵌套树,也不是嵌入相似度——就是 hash(str(observation))。这一切建立在一个物理假设上:多条 rollout 会撞到同一个中间状态。撞不到,方法退化。

GiGPO 论文里把方法叫做 “group-in-group”,读起来像嵌套的树结构、像 MCTS 的分层探索。真读 gigpo/core_gigpo.py,你会发现所谓的 “group-in-group” 就是两句代码:scores = episode_advantages + step_advantage_w * step_advantagescore_gigpo.py:183)。所谓的 “anchor state matching” 就是 hash(str(observation))——精确串匹配,或者退一步走 difflib.SequenceMatcher.ratio() >= 0.95 的模糊匹配。

这不是”包装”批评,是理解 GiGPO 能不能用在你的场景的关键。这个方法 arXiv:2505.10978(NeurIPS 2025)在 ALFWorld 上比 GRPO 高 12 个百分点、在 WebShop 上高 9 个百分点、几乎零额外开销——但这个红利依赖一个物理假设:多条 rollout 会撞到同一个中间状态。撞不到,A_step 退化成 0,GiGPO 就是 GRPO。

我这篇从 langfengQ/verl-agent 的源码入手,回答一个问题:在你的 agent 训练场景里,GiGPO 会给你梯度,还是给你噪声?

名词速查

术语一句话解释
GRPOGroup Relative Policy Optimization。同一 prompt 采样 N 条轨迹,用组内 mean/std 把 return 归一化当 advantage。critic-free、省显存。
credit assignment长轨迹里,reward 到底该归功于哪一步。50 步的 episode 只有终态 pass/fail,中间 49 步全按同一个 advantage 更新,等于没做 credit assignment。
episode-level advantage整条轨迹的 return 减去组内平均,度量”这条整体跑得好不好”。GRPO 用的就是这个。
step-level advantage同一 anchor state 出发的多条轨迹之间对比,度量”从这个岔路做出的这一步好不好”。GiGPO 新加的。
anchor state用 observation 的哈希值当 key,把在中途撞到同一个 observation 的多条 rollout 聚起来当对照组。
joint advantageGiGPO 的最终 advantage:A = A_episode + w * A_step,默认 w=1加法,不是嵌套。

如果对 GRPO 和 RLVR 那条主线没底,先读 DeepSeek-R1 学习笔记——GRPO 是那篇的主角,本篇是它长成 agent 版之后的样子。

根本约束:状态重访率决定一切

写下一句可反驳的话:GiGPO 之所以长成”两个 group 相加”这个形状,是因为 agent 训练里的 credit assignment 需要一个可以做局部对照的岔路点——而只有当多条 rollout 撞到同一个 observation 时,这个岔路点才存在。

这是可反驳的:如果每条轨迹的中间状态都独一无二,anchor 组只有单人,A_step 会怎么样?读 core_gigpo.py:step_norm_reward 第 366-375 行:

if len(id2score[idx]) == 1:
    id2mean[idx] = torch.mean(torch.tensor(id2score[idx]))  # = 这个 singleton 的分数
    id2std[idx] = torch.tensor(1.0)
...
if remove_std:
    scores[i] = scores[i] - id2mean[index[i]]  # score - score = 0

单人 anchor 组,A_step = 0。 这条轨迹的这一步 advantage 就退化成纯 A_episode——也就是 GRPO。论文自己也写了这一点:“in the extreme case where no states are repeated (A^S=0), it naturally degrades to GRPO.”

所以 GiGPO 相对 GRPO 的净收益完全由一个数决定:有多少比例的 step 落在 size >= 2 的 anchor 组里? 这个比例低于某个阈值,你付了 anchor 计算的开销(≤1%)却拿不到 credit assignment 的红利。

崩点:换成最朴素的 GRPO 会先在哪里崩? GRPO 把整条轨迹的 return 平均分给每一步 token。ALFWorld 一条 episode 有 50 步、20k+ token(论文 §4.1),只有最后一步有 pass/fail 信号——前 49 步的关键决策和后续无关重复动作拿到完全相同的梯度信号。这不是 credit assignment 差,是没有 credit assignment。GiGPO 的补丁:把中途的岔路点标出来,让”从这个岔路选 A 的那一队”和”选 B 的那一队”能在局部对照里互相当基线。

共同祖先:这个想法有更老的名字——

  • MCTS 里的 UCB1 也是”同一节点下的孩子做对照”,但那是树、有 backprop;GiGPO 是 flat 的哈希桶。
  • 反事实推理(counterfactual advantage)在 multi-agent RL 里已经存在,GiGPO 是它的廉价近似版:不用 counterfactual policy,直接借”另一条真跑过的轨迹”当反事实。
  • Hindsight Experience Replay 也用”事后重看轨迹”的思路——GiGPO 把”事后”缩短到”同一批 rollout 内部”。

叫它 “group-in-group” 会让人以为是嵌套结构。真实结构是:两个独立的分组各算一次 mean-normalized advantage,然后加起来。加法。

主结构:一张伪代码读完

不动源码语义地压缩 compute_gigpo_outcome_advantagecore_gigpo.py:158-183):

# 伪代码,对应 core_gigpo.py:158; batch/mask/dtype 省略
def compute_gigpo_advantage(rewards, anchor_obs, prompt_idx, traj_uid, w=1.0):
    # 1) episode 组:同一 prompt 下所有 rollout 的 return 归一化
    A_episode = episode_norm_reward(rewards, prompt_idx, traj_uid)

    # 2) 锚点分组:按 (prompt, hash(obs)) 把 step 分到同一桶
    #    单人桶保留但不产生信号 (A_step=0)
    step_group = build_step_group(anchor_obs, prompt_idx)

    # 3) step 组:同一锚点桶内的 discounted return 归一化
    A_step = step_norm_reward(step_rewards, step_group)

    # 4) 相加得到 joint advantage
    return A_episode + w * A_step

第 2 步的哈希函数就是 to_hashablecore_gigpo.py:33-46):字符串/数字直接用,np.array 转 tuple(flatten()),dict 转 tuple(sorted(items))。没有嵌入、没有语义相似度。默认 enable_similarity=False 走精确匹配;打开后走 difflib.SequenceMatcher.ratio() >= 0.95core_gigpo.py:71-83)——这是字符串编辑距离的近似,不是语义匹配。

第 3 步的 step reward 是 compute_step_discounted_returnscore_gigpo.py:85-135):从轨迹末尾往前扫,R_t = r_t + γ · R_{t+1}。所以 anchor 处的 “step reward” 其实是从这一步开始到 episode 结束的折扣回报——这才是 credit assignment 的信号载体。

状态重访率:亲手算一次

论文没有报告”平均 anchor 组大小”这个统计,但源码里有——build_step_group 结尾打印 Avg size of step-level groupcore_gigpo.py:346),训练时会滚在日志里。

我写了个 40 行 Python 脚本(标准库,无外部依赖)模拟这个统计:假设 n_traj 条 rollout、每条 ep_len 步、状态从大小为 state_space 的池子里 IID 抽取,看有多少比例的 step 落在 size >= 2 的桶里。

import random
from collections import Counter

def simulate(n_traj, ep_len, state_space, seed=0):
    random.seed(seed)
    all_states = [random.randint(0, state_space - 1)
                  for _ in range(n_traj * ep_len)]
    counter = Counter(all_states)
    group_sizes = [counter[s] for s in all_states]
    return (sum(group_sizes) / len(group_sizes),                       # 平均组大小
            sum(1 for k in group_sizes if k >= 2) / len(all_states))   # 可用 %

for ss in (20, 100, 200, 500, 1000):
    avg, usable = simulate(8, 50, ss)
    print(f"state_space={ss:>4}  avg_group={avg:5.2f}  usable={usable:.1%}")

真实输出:

state_space=  20  avg_group=20.64  usable=100.0%
state_space= 100  avg_group= 5.04  usable= 98.2%
state_space= 200  avg_group= 3.29  usable= 86.5%
state_space= 500  avg_group= 1.90  usable= 58.2%
state_space=1000  avg_group= 1.46  usable= 35.5%

读法:state_space 是”这一批 rollout 能撞到的不同 observation 总数”的上界

  • state_space=20(比如 ALFWorld 的一个小房间,agent 反复走进走出):平均 anchor 组 20 人,100% step 有对照——GiGPO 收益最大化。
  • state_space=200(复杂室内环境):平均组 3.3 人,86% step 可用——GiGPO 明显有增量。
  • state_space=1000(可想象为 WebShop 的中等 SKU 目录,或代码补全时不同的文件光标位置):平均组 1.5 人,只有 35% step 能拿到 step 信号——GiGPO 相对 GRPO 的净收益要打折
  • 一个真实浏览器的 DOM:每次页面渲染的字符串都可能微妙不同,effective state space 可能上万——enable_similarity=True 才有戏,但 SequenceMatcher 是 O(n²) 的字符串比较,长 observation 上会变慢,且 first-fit 聚类对顺序敏感。

这就是”论文没做的一处消融”——GiGPO 的收益曲线相对状态重访率呈现什么样的形状。ALFWorld 的高收益是因为它落在”重访率极高”的这一端;把方法搬到状态重访率低的场景,你需要自己先测这个数。

三处源码值得读的细节

细节一:episode 归一化的 mean 跨 step 计算。 episode_norm_rewardcore_gigpo.py:216)默认 compute_mean_std_cross_steps=True——同一 prompt 下所有 rollout 的每个 step 都参与算 mean,而不是”每条轨迹一个 return,N 条一个均值”。这让样本量从 n_traj 变成 n_traj * ep_len,方差小很多。这是稳定训练的一个非平凡工程决定。

细节二:joint advantage 用加法,不用乘法或门控。 scores = episode_advantages + step_advantage_w * step_advantages(第 183 行)。默认 w=1。加法有个隐含语义:episode 好但 step 坏、episode 坏但 step 好都能被表达——两个信号独立线性叠加,谁大谁主导。用乘法会引入”两个都必须为正才更新”的耦合。用门控会引入超参。加法最简单且效果好,是论文的关键工程直觉。

细节三:mode=“mean_norm” 是代码默认,不是论文表格里的 “w/ std” 版本。 第 165 行 mode: str = "mean_norm"——即 remove_std=True,只做减均值、不除 std。论文 Table 里 “GiGPO w/o std” 通常和 “w/ std” 打平或略高(ALFWorld 7B: 90.2 vs 90.8;WebShop 7B: 86.2 vs 论文未列同规格 “w/ std”)。代码默认的省 std 模式,减少了组内方差不稳带来的震荡——尤其对单人组,std=1 是默认值,除法退化。省掉 std 除法反而更稳。这是一个不引人注意但重要的工程简化。

论文数字:GiGPO vs GRPO

论文 Table 1、2、3(Qwen2.5-Instruct 系列,我从 arXiv HTML 版核对,未通读 PDF):

基准模型GRPOGiGPO w/o stdGiGPO w/ std
ALFWorld (成功率)1.5B72.8±3.686.1±4.786.7±1.7
ALFWorld (成功率)7B77.6±5.290.2±2.390.8±1.3
WebShop (score)1.5B75.8±3.583.5±1.8
WebShop (score)7B79.3±2.886.2±2.6
WebShop (成功率)7B66.1±3.775.2±3.8

几个观察:

  1. 方差降得比均值升得更整齐(ALFWorld 7B 的 std 从 5.2 降到 1.3)。这是 credit assignment 变精细的直接表现——梯度里”信号 vs 噪声”的比值改善。
  2. 7B 相对 1.5B 的绝对增益更小(ALFWorld 7B 提 13.2 分 vs 1.5B 提 13.9 分)——但方差改善更明显。小模型吃 credit assignment 红利更急。
  3. 计算开销可忽略:anchor 分组 0.01s/迭代,step 归一化 0.53s,占总迭代时间 <0.002%(论文报告的 ~363s)。这个方法的”性价比”极高。

适用边界:什么时候不要用 GiGPO

我倾向于(这是判断,第二档证据,由前面的机制分析支撑)GiGPO 在以下四种情况会打折或失效:

  1. 状态重访率低于 50%:先用上面的 simulate 脚本估你的场景的 usable %,或者训练时打开 summarize_group_size 看真实分布。低于 50% 意味着一半以上 step 拿到 A_step=0,收益打对折。
  2. Observation 是自由文本 + 长上下文:DOM 快照、完整代码文件、长对话历史——精确哈希几乎撞不到,enable_similarity 又慢又粗。这种场景下考虑先做”状态压缩”——只 hash observation 的关键子集(当前 URL、当前光标位置、当前 stack frame)而不是全文。
  3. 短 episode(<10 步):credit assignment 问题本身不严重,GRPO 已经够用。GiGPO 的复杂度花在”识别岔路点”上,短 episode 没几个岔路点。
  4. 稀疏动作空间 + 稠密状态空间:agent 可能的动作很少(比如 5 个),但每个动作会触发 observation 大幅变化。这种情况 anchor 桶很难聚起来。

小结

  • 根本约束:agent 训练的 credit assignment 需要局部对照的岔路点;这个岔路点只在多条 rollout 撞到同一 observation 时存在。GiGPO 的红利完全由”状态重访率”决定。
  • 崩点:换成 GRPO,50 步 episode 的前 49 步全按同一个 advantage 更新——不是 credit assignment 差,是没有。
  • 共同祖先:MCTS 的同节点孩子对比 + counterfactual advantage 的廉价近似 + HER 的事后重看。“group-in-group” 是加法,不是嵌套。
  • 可带走的动作:如果你在训 agent 且用了 GRPO,先跑一次 Counter 看 observation 的重访率,就能预判 GiGPO 值不值得接。ALFWorld 是”高重访率”这一端的典型;网页浏览是”低重访率”这一端的典型。

五分钟实验:你的场景吃不吃 GiGPO

在你已有的 agent rollout 日志上(哪怕是 GRPO 训过程中的 log 抽样),跑这段:

from collections import Counter

# 从你的 rollout 里抽 100+ 步的 observation 列表
# 每条 observation 是 agent 那一步看到的文本/字典/截图 caption
observations = load_observations_from_your_rollout_logs()

def state_repeat_stats(obs_list):
    # 用 GiGPO 的 to_hashable 方式做键:字符串直接用,字典排序 items
    keys = [
        s if isinstance(s, str) else tuple(sorted(s.items()))
        for s in obs_list
    ]
    c = Counter(keys)
    sizes = [c[k] for k in keys]
    return {
        "avg_group_size": sum(sizes) / len(sizes),
        "usable_pct": sum(1 for k in sizes if k >= 2) / len(sizes),
        "singleton_pct": sum(1 for k in sizes if k == 1) / len(sizes),
    }

print(state_repeat_stats(observations))

判据(我根据论文和源码机制推导,未经大规模消融验证):

  • usable_pct > 80%:直接上 GiGPO,稳收益。
  • 50% < usable_pct < 80%:可以试,但预期增益打七折。
  • usable_pct < 50%:先做状态压缩——挑一个 observation 的低维摘要子集(URL 前缀、当前 tab 名、action 上一步)当 anchor,而不是整个 observation。压缩后重跑上面的统计再决定。
  • usable_pct < 20%:GiGPO 大概率退化成 GRPO,别浪费实验。

再进一步的话,直接读 gigpo/core_gigpo.py 从第 158 行到第 350 行——不到 200 行,是我近期读过的 RL 算法实现里最好读的一份。读完再回论文,你会知道哪些是”论文措辞的漂亮包装”,哪些是”代码里真做的机械动作”。

参考来源

论文

源码

  • langfengQ/verl-agent — GiGPO 官方实现(veRL 的扩展)。本文所有源码定位均来自 main 分支,具体到:gigpo/core_gigpo.py:33-46(哈希函数)、:71-83(相似度匹配)、:85-135(折扣回报)、:158-183(joint advantage 入口)、:216-274(episode 归一化)、:277-350(anchor 分组)、:353-395(step 归一化)。

本站相关