PPO 中 GAE 的分 chunk 并行计算(基于 slime 的实现) 对应知乎原文:《PPO 中 GAE 的分 chunk 并行计算》(https://zhuanlan.zhihu.com/p/1975237289425798560) 对应代码: PR #850 — Chunk-Scan GAE TL;DR 这篇文章中,作者围绕 slime 框架里的 PPO + GAE 做了一次性能改造: 背景:在 agentic RL 场景里,在序列超长的时候,slime 原本的 GAE 计算是按 sample 分批串行从尾部到头扫一遍。而这将直接变成训练瓶颈。
对应知乎原文:《PPO 中 GAE 的分 chunk 并行计算》(https://zhuanlan.zhihu.com/p/1975237289425798560)
对应代码:THUDM/slimePR #850 — Chunk-Scan GAE
这篇文章中,作者围绕 slime 框架里的 PPO + GAE 做了一次性能改造:
背景:在 agentic RL 场景里,在序列超长的时候,slime 原本的 GAE 计算是按 sample 分批串行从尾部到头扫一遍。而这将直接变成训练瓶颈。
做的事:
lastgaelam 逐步推出最后的 GAE。效果:
chunk_size,在不 OOM 的前提下,chunk_size 越大加速越明显。在 RLHF / Agentic RL 里,PPO 仍是一个非常常用、表现稳定的算法。我们需要在每个 token 上计算 advantage,最常见的就是 GAE(Generalized Advantage Estimation)。而 GAE 的标准写法是一个从后往前的递推公式,对序列长度 T 来说,是 O(T) 串行依赖。
在 slime 中,GAE 的算法实现如下:
lastgaelam = torch.zeros(B, device=device, dtype=dtype) adv_rev = [] for t in reversed(range(max_len)): next_value = full_values[:, t + 1] if t < max_len - 1 else 0.0 delta = full_rewards[:, t] + gamma * next_value - full_values[:, t] lastgaelam = delta + gamma * lambd * lastgaelam adv_rev.append(lastgaelam) full_advantages = torch.stack(adv_rev[::-1], dim=1) # [B, max_len]
slime 在一开始的实现里,追求的是“支持变长序列”,优点是在模型计算时,不需要 padding 到所有序列的 max_len, 避免浪费无效的计算,所以在计算 GAE 时,是一个序列一个序列计算而不是拼成 batch 计算,造成了性能瓶颈,我们很快改成常见的 “padding 到 max_len,再按 batch 计算 GAE” 的写法。但是可惜的是,这并不足以达到可能的最佳性能,它在时间维度仍是串行,而这导致了在长序列场景下依然很吃力。
在此基础上,作者结合 @sonta 在讲 linear attention 时提到的 “分 chunk 并行 + chunk 间轻量递推” 的思路,尝试把 GAE 也改造成一个 chunk 级别可并行的“前缀扫描(scan)”问题。
想对 GAE 并行,其实有一个非常优雅的方案,直接写成矩阵乘法:
把 GAE 写成 A_t = \sum_{k=t}^{T-1} w^{k-t} \delta_k(其中 w = \gamma \lambda)
构造一个 T×T 的上三角权重矩阵 W,然后做 A = \delta W^\top
这是完全可以并行的,但是这直接导致了时间复杂度和空间复杂度都是 O(T²)。一旦 T 达到了 64K、128K 的级别,会直接 OOM。
torchrl 中有使用 conv1d 将时间复杂度降到 O(T) 的方案,但是空间复杂度依然是 O(T²),因此还是会有上面这个 OOM 问题。
因此,我们希望找到一个能同时兼顾并行度,又能保证显存可控的 GAE 计算方式。
在开始前,我们先回顾一下标准的 GAE:
我们记 delta 为
则 GAE 的 advantage 为
也可以写成后向递推的形式:
slime 目前的版本其实已经给出了答案:
lastgaelam = torch.zeros(B, device=device, dtype=dtype) adv_rev = [] for t in reversed(range(max_len)): next_value = full_values[:, t + 1] if t < max_len - 1 else 0.0 delta = full_rewards[:, t] + gamma * next_value - full_values[:, t] lastgaelam = delta + gamma * lambd * lastgaelam adv_rev.append(lastgaelam) full_advantages = torch.stack(adv_rev[::-1], dim=1) # [B, max_len]
优点:实现简单,数值稳定;
缺点:这个版本在时间维度完全串行,长序列下性能不行。
利用前向展开式:
我们可以构造一个 T×T 的权重矩阵 W:
于是有:
我们可以把整条序列拆成若干个长度为 C 的 chunk:
第一个 chunk:0 - C-1 第二个 chunk:C - 2C-1 ... 第 c 个 chunk:cC - (cC + L_c - 1)
在反向序列上定义 GAE 递推:
对于第 c 个 chunk,定义“跨 chunk 状态”:
c = 0 时,有 s_{\text{prev}} = S_{-1} = 0;
现在考虑 chunk c 内部的第 t 个元素(局部索引 t = 0..L_c-1):
把“当前 chunk 内”的部分单独拿出来:
于是最终公式可以写成:
这意味着:
s_prev,串行递推即可。时间复杂度:O(T·C)
空间复杂度:O(T + C²)
上面是非常严谨的公式推导过程,其实简单来说,Chunk-scan 的核心想法就是:
THUDM/slime PR #850 — Chunk-Scan GAE 实现的伪代码以下是一个展示如何把 Chunk-Scan GAE 写成批量计算函数的伪代码,代码原型来自:THUDM/slime PR #850 — Chunk-Scan GAE
function chunked_gae(rewards, values, gamma, lambda, chunk_size): w = gamma * lambda # 1. 计算每一步的 δ_t deltas = compute_deltas(rewards, values) # δ_t = r_t + γV_{t+1} - V_t # 2. 反向时间顺序(从后往前的递推 -> 在反向序列上从左往右) deltas_rev = reverse_time(deltas) # 3. pad 到 chunk_size 的整数倍,并拆成若干个 chunks deltas_chunks = split_into_chunks(deltas_rev, chunk_size) # 4. 为“每个 chunk 内部”的扫描预计算一个小核: # 给定一段 Δ[0..C-1],算出 s_local[t] = Σ_{k≤t} w^(t-k) * Δ[k] kernel = build_chunk_kernel(chunk_size, w) # C×C 的上三角矩阵 pow_vec = build_power_vector(chunk_size, w) # [w^1, w^2, ..., w^C] # 5. 所有 chunk 内部并行做局部 scan # local_scan[c, t] = s_local^(c)[t] local_scans = [] for each chunk in deltas_chunks in parallel: s_local = chunk @ kernel # 这里用任意并行实现都行 local_scans.append(s_local) # 6. 在 chunk 之间串行传播“前缀状态” s_prev s_prev = 0 full_scan_rev = empty_like(deltas_rev) for c from 0 to num_chunks-1: s_local = local_scans[c] # 当前 chunk 内部的结果,长度 L_c # 注入跨 chunk 的状态: # S_global[t] = s_local[t] + w^(t+1) * s_prev S_global = s_local + s_prev * pow_vec[0:L_c] write_into(full_scan_rev, chunk_index=c, values=S_global) # 下一个 chunk 的起点状态 = 当前 chunk 最后一个位置 s_prev = S_global[L_c - 1] # 7. 去掉 padding,反向回正向时间 advantages = reverse_time(remove_padding(full_scan_rev)) # 8. returns 一般就是 V_t + A_t returns = values + advantages return advantages, returns
根据原文的实验结果,实现效果非常可观:
| No chunk | chunk size = 64 | chunk size = 128 | chunk size = 256 | |
|---|---|---|---|---|
| B=256, T=131072 | 5.935994s | 0.070122s | 0.034059s | 0.018390s ( x317 ) |
| B=128, T=65536 | 2.902570s | 0.232986s | 0.017645s | 0.009134s |
可以看到:
chunk_size=256 时,加速比约 317×;chunk_size=256 时,加速比也非常可观;只要有足够显存能用来提升 chunk size,并行度就能大幅度增加,GAE 的计算时间也能相当可观地被缩减。
Chunk-Scan 已被作为默认的训练行为,因此对于用户的安装或迁移,仅需更新镜像即可。
THUDM/slime PR #850 — Chunk-Scan GAE 的改动。截止至 11/24,官方还未更新 docker 镜像。THUDM/slime PR #850 — Chunk-Scan GAE 的改动。截止至 11/24,官方还未更新 docker 镜像。更系统的 benchmark & 可视化工具
提供一键脚本,方便用户评估自己任务是否值得开启 Chunk-Scan。
更全面地测试整体框架的性能,更细粒度地测量各个部分的耗时情况,找出类似的潜在的问题。
检查其他部分的代码是否也存在可以通过修改算法提升并发度的情况,如果有,需要探索优化的可能性。