- 文档
- 教程
- 人工智能
- 大模型
- RLHF
【免费下载链接】Awesome-ML-SYS-Tutorial
My learning notes for ML SYS.
本篇技术指南围绕开源仓库 rlhf/slime/batch-GAE/ppo-gae-chunk.md 所记录的 slime 框架 PPO 训练管线改造展开:面对 agentic RL 长序列场景下 GAE(Generalized Advantage Estimation)串行递推成为训练瓶颈的问题,借鉴 linear attention 的"分 chunk 并行 + chunk 间轻量递推"思路,将 GAE 改造成 chunk 级可并行的前缀扫描(prefix scan)问题。读完本文,你将掌握 GAE 串行解、纯矩阵解、Chunk-Scan 解三种方案的原理与取舍,能够复现 100×–300× 的 GAE 计算加速,并理解如何在类似 RL 框架中排查与消除训练流水线中的串行瓶颈。
1. TL;DR
这篇文章围绕 slime 框架里的 PPO + GAE 做了一次性能改造,核心结论如下:
- 背景:在 agentic RL 场景里,序列超长时,slime 原本的 GAE 计算是按 sample 分批、从尾部到头串行扫描一遍,这直接变成训练瓶颈。
- 做法:
- slime 先把传统的串行后向递推计算 GAE 的方式,修改为先将 GAE 按时间分为多个 Chunk,之后按时间逆序,用当前遍历时间点的 Chunk 和上一个 Chunk 计算出的
lastgaelam逐步推出最后的 GAE。 - 本文则进一步借鉴linear attention的"分 chunk 并行 + chunk 间轻量递推"思路,通过分块前缀扫描,将多个计算完成的局部 GAE 合并,计算出最终 GAE。改造后每个 Chunk 的计算之间不再有依赖,可以充分并行。
- slime 先把传统的串行后向递推计算 GAE 的方式,修改为先将 GAE 按时间分为多个 Chunk,之后按时间逆序,用当前遍历时间点的 Chunk 和上一个 Chunk 计算出的
- 效果:
- 在 slime 中,GAE 计算时间得到100×–300×的加速(实测最大约317×);
- 并行度取决于
chunk_size,在不 OOM 的前提下,chunk_size越大加速越明显。
需要说明:本文依据仓库中的学习笔记整理,涉及的外部技术报告与上游 PR 细节以仓库记录为准。slime 框架本身是"基于 SGLang 与 Megatron LM 作为唯一后端"的 RL 训练框架,其整体架构与数据流可参考仓库中 rlhf/slime/code-walk-through/readme.md 的代码走读记录。
2. 技术背景:为什么要搞 GAE 的 Chunk-Scan?
2.1 为什么现在 GAE 会变成瓶颈?
在 RLHF / Agentic RL 里,PPO 仍是一个非常常用、表现稳定的算法。我们需要在每个 token 上计算 advantage,最常见的就是 GAE。而 GAE 的标准写法是一个从后往前的递推公式:对序列长度 T 来说,是 O(T) 的串行依赖——A_t依赖A_{t+1},无法直接并行。
在 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]其中gamma是折扣因子,lambd是 GAE 的 λ 参数,full_values是 critic/value 网络输出的每个 token 的状态价值估计,full_rewards是每个 token 的奖励序列。
slime 在一开始的实现里追求"支持变长序列",优点是在模型计算时不需要 padding 到所有序列的 max_len,避免浪费无效计算;代价是计算 GAE 时一个序列一个序列计算、而不是拼成 batch 计算,造成了性能瓶颈。后来很快改成常见的"padding 到 max_len,再按 batch 计算 GAE"的写法——但这并不足以达到可能的最佳性能:它在时间维度上仍然是串行的,在长序列场景下依然很吃力。
在此基础上,作者结合 linear attention 中"分 chunk 并行 + chunk 间轻量递推"的思路,尝试把 GAE 也改造成一个 chunk 级别可并行的"前缀扫描(scan)"问题。
2.2 完全矩阵化计算下的爆显存问题
想对 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 计算方式——这正是 Chunk-Scan 要解决的问题。
3. 架构设计:从串行 GAE 到 Chunk-Scan GAE
3.1 标准 GAE 回顾
在开始前,先回顾一下标准的 GAE。记 delta 为:
$$ \delta_t = r_t + \gamma V_{t+1} - V_t $$
则 GAE 的 advantage 为:
$$ A_t = \sum_{k=t}^{T-1} (\gamma \lambda)^{k-t} \delta_k $$
也可以写成后向递推的形式:
$$ A_t = \delta_t + \gamma \lambda A_{t+1}, \quad t = T-1, T-2, \dots, 0 $$
可以看到,后向递推形式天然是串行的:每个时间步都要等它后面一个时间步的结果。三种解法的差异,本质上就是"如何打破这条串行依赖链"。
3.2 方案一:串行解法
slime 目前的版本给出的答案(与 2.1 节同一段代码):
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]- 优点:实现简单,数值稳定;
- 缺点:这个版本在时间维度完全串行,长序列下性能不行。
值得注意的细节是next_value的边界处理:当t == max_len - 1(最后一个 token)时没有V_{t+1},按 0.0 处理;adv_rev按逆序收集结果后,通过torch.stack(adv_rev[::-1], dim=1)再翻转为正向时间顺序,得到形状为[B, max_len]的完整 advantage 矩阵。这个边界约定在后续所有并行化方案中都必须保持一致。
3.3 方案二:纯矩阵解法
利用前向展开式:
$$ A_t = \sum_{k=t}^{T-1} w^{k-t} \delta_k,\quad w = \gamma \lambda $$
构造一个 T×T 的权重矩阵 W:
$$ W_{t,k} = \begin{cases} w^{k-t}, & k \ge t \ 0, & k < t \end{cases} $$
于是有:
$$ A = \delta W^\top $$
- 优点:矩阵乘法可以在 GPU 上高度并行;
- 缺点:非常容易 OOM(时间、空间复杂度均为 O(T²),T 到 64K/128K 量级直接爆显存)。
这个方案的意义在于指明了"GAE 是可以并行计算的"这一方向,但直接矩阵化的代价太高,需要一个中间路线。
3.4 方案三:Chunk-Scan(分块前缀扫描)
我们可以把整条序列拆成若干个长度为 C 的 chunk:
第一个 chunk:0 - C-1 第二个 chunk:C - 2C-1 ... 第 c 个 chunk:cC - (cC + L_c - 1)在反向序列上定义 GAE 递推:
$$ S_i = \widetilde{\delta}i + w S{i-1}, \quad w = \gamma \lambda, \quad S_{-1} = 0 $$
其中 $\widetilde{\delta}_i$ 表示反向时间序列上第 i 个位置的 delta(即把正向的 $\delta$ 倒序后,从左往右做递推,就等价于原始从后往前的递推)。对于第 c 个 chunk,定义"跨 chunk 状态":
$$ s_{\text{prev}} = S_{cC - 1} $$
c = 0 时,有 $s_{\text{prev}} = S_{-1} = 0$。
现在考虑 chunk c 内部的第 t 个元素(局部索引 t = 0..L_c-1):
- 全局索引 $i = cC + t$;
- 展开递推关系得到:
$$ \begin{aligned} S_{cC + t} &= \widetilde{\delta}_{cC + t}
- w \widetilde{\delta}_{cC + t - 1}
- \cdots
- w^t \widetilde{\delta}_{cC}
- w^{t+1} S_{cC - 1} \end{aligned} $$
把"当前 chunk 内"的部分单独拿出来:
$$ s^{(c)}t = \sum{k=0}^{t} w^{t-k} \widetilde{\delta}_{cC + k} $$
于是最终公式可以写成:
$$ \boxed{ S_{cC + t} = s^{(c)}t + w^{t+1} , s{\text{prev}}, \quad t = 0, \dots, L_c - 1 } $$
这意味着:
- 局部部分$s^{(c)}_t$ 可以在 chunk 内用矩阵/conv 并行算——它只依赖当前 chunk 内部的 $\widetilde{\delta}$,是 chunk 内部的前缀加权和,形式上与线性 attention 的 chunk 内扫描完全同构;
- 跨 chunk只需要维护一个标量状态
s_prev,chunk 之间串行递推即可,每个 chunk 只需向后续 chunk 传递最后一个位置的全局状态 $S_{cC + L_c - 1}$。
时间复杂度:O(T·C)(chunk 内矩阵乘 + chunk 间线性扫描)空间复杂度:O(T + C²)(存整条序列的中间结果 + chunk 内的 C×C 核矩阵)
与 O(T²) 的纯矩阵解相比,显存从二次方降为线性,chunk 数(T/C)越少,并行度越高。
简单来说,Chunk-Scan 的核心想法就是三步:
- 切分:把长序列切成若干小 chunk;
- 并行局部扫描:让 GPU 并行计算每个 chunk 内的递推,上面的公式得出了可以并行计算的部分 $s^{(c)}_t$;
- 合并:再把这些 chunk 的结果通过标量状态
s_prev组合起来,还原全局 GAE。
这正是 linear attention(如 chunked prefix scan 变体)里"chunk 内并行、chunk 间递推"思想向 GAE 的一次迁移:GAE 的递推 $S_i = \widetilde{\delta}i + w S{i-1}$ 与一阶自回归形式共享同一类可结合算子结构,因此可以套用同样的并行扫描技术。
3.5 Chunk-Scan GAE 的实现伪代码
以下是展示如何把 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对这段伪代码的几个实现要点展开说明:
kernel(C×C 上三角矩阵):其元素为 $w^{t-k}$($t \ge k$,否则为 0),正是 3.3 节矩阵解中 W 的一个局部子块。chunk @ kernel一行即可算出 chunk 内所有位置的局部前缀加权和 $s^{(c)}_t$,等价于一次小的矩阵乘法,GPU 高度并行。pow_vec(幂向量):预先算好 $[w^1, w^2, \dots, w^C]$,用于把s_prev传播到 chunk 内每个位置(第 t 个位置需要乘 $w^{t+1}$)。注意 t 从 0 开始,因此下标偏移为 1。s_prev的更新:只需要取当前 chunk 的最后一个全局状态S_global[L_c - 1]。由于 chunk 内所有位置的状态可以一次性算出,s_prev的传播只需要在 chunk 粒度上进行(共 T/C 次标量运算),这就是"chunk 间轻量递推"的含义。- padding 与边界:
split_into_chunks之前需要把deltas_revpad 到 chunk_size 的整数倍;最终remove_padding去掉这部分,保证输出与原始序列严格对齐,数值结果与串行解法一致。 returns = values + advantages:这是 PPO 训练中计算 policy loss 与 value loss 的标准一步,说明该函数直接产出训练管线可用的数据。
在 slime 的 PPO 训练管线中,这条计算链位于训练侧(Training/Megatron 后端):rollout 阶段(SGLang 后端)负责生成 token 序列与 reward,进入训练阶段后,critic 网络给出 value 估计,随后执行上述 GAE 计算得到每个 token 的 advantage,再进入策略与价值函数的更新。仓库中 rlhf/slime/code-walk-through/readme.md 记录了 slime"Training (Megatron) + Rollout (SGLang) + Data Buffer"的分离式架构,可以帮助理解 GAE 计算在整个流水线中的位置。
4. 实现效果
根据仓库笔记记录的实验结果,实现效果非常可观:
| 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 |
可以看到:
- 在T=131072(128K 超长序列)、chunk_size=256时,加速比约317×——从接近 6 秒压到不足 20 毫秒;
- 在T=65536、chunk_size=256时,加速比同样非常可观——从 2.9 秒压到约 9 毫秒。
结论:只要有足够显存来提升 chunk size,并行度就能大幅增加,GAE 的计算时间也能被相当可观地缩减。这里也解释了 chunk_size 的调参方向:chunk_size 越大,chunk 数越少,chunk 间的串行递推步数越少、并行度越高,但 chunk 内的 C×C 核矩阵(空间复杂度 O(C²))占用显存也越大,因此需要在显存预算内尽量增大 chunk_size。作为参考:表中chunk_size=64与chunk_size=128、chunk_size=256的耗时递减趋势与这一规律一致。
5. 具体使用方法:在 slime 里怎么用 Chunk-Scan GAE?
Chunk-Scan 已被作为默认的训练行为,因此对于用户的安装或迁移,仅需更新镜像即可,无需修改任何训练参数。
5.1 安装
- 拉取当前最新版本 docker 镜像,确保其包含 Chunk-Scan GAE 的改动(仓库笔记记录:截止至 11/24,官方 docker 镜像尚未更新该改动,需注意发布时间线);
- 根据 slime 官方指引部署服务(
docs/en/get_started/quick_start.md,仓库内可参考 rlhf/slime/code-walk-through/readme.md 中记录的框架结构了解部署形态); - 使用默认参数即可使用 Chunk-Scan 训练 PPO——因为该功能是作为默认行为合入的,不需要显式开启。
5.2 迁移
- 升级至当前最新版本 docker 镜像,确保其包含 Chunk-Scan GAE 的改动;
- 恢复训练即可(无需修改既有 checkpoint 或训练配置,计算结果与串行版在数值上等价)。
6. 未来计划
仓库笔记中记录了作者对后续工作的规划,体现了"性能优化是系统性工程"的思路:
- 更系统的 benchmark 与可视化工具:提供一键脚本,方便用户评估自己的任务是否值得开启 Chunk-Scan(即判断 GAE 是否确实是当前训练管线的瓶颈);
- 更全面地测试整体框架的性能:更细粒度地测量各个部分的耗时情况,找出类似的潜在问题;
- 检查其他部分的代码:排查是否还有其他通过修改算法提升并发度的机会,若有,探索优化的可能性。
这三点实际上给出了一套通用的"瓶颈排查—优化—验证"工作流:先细粒度 profiling 定位串行热点,再针对热点做算法级并行化改造,最后用工具化手段让优化决策可复现、可推广。
7. 工程附录:踩过的坑 & 学到的东西
"用实验结果纠正工程直觉":GAE 变成瓶颈这件事本身,就是一个典型案例。在实验数据真正跑出之前,很难想到 GAE 计算会成为 PPO 流水线的瓶颈——直觉上它只是"一个 O(T) 的循环",但长序列 + 逐 sample 串行 + 训练步数放大后,累积耗时非常可观。因此,对一个成熟的框架来说,应该把性能测试的粒度划分得足够细,从而发现一些设计之初可能会忽视的问题。
并行化的三层递进:从"逐 sample 串行"到"batch 内时间维串行"再到"Chunk-Scan 全并行",每一步都以"不改变数值语义"为前提。这个改造路径对任何 RL/序列算法都适用:先保证正确性,再逐步把依赖链从数据维度转移到可并行的计算结构上。
复杂度取舍是核心权衡:纯矩阵解 O(T²) 虽然完美并行却必然 OOM,Chunk-Scan 通过把"时间维并行"降级为"chunk 内并行 + chunk 间标量递推",把空间复杂度压到 O(T + C²),在显存与并行度之间找到了工程上可用的平衡点。这也是 linear attention 类方法(chunked scan)能够落地的根本原因——并行的代价是显存,而显存是可以由 chunk_size 显式调节的。
如果你正在使用 slime 或其他 RLHF 框架训练超长上下文(64K/128K 级别)的 agentic 任务,建议按照本文第 6 节的思路先做一次细粒度 profiling:若 GAE 计算在训练 step 中的占比显著,Chunk-Scan 方案(以及配套的 chunk_size 调参)就是一项低成本、高收益、且数值等价的改造。
- 文档
- 教程
- 人工智能
- 大模型
- RLHF
【免费下载链接】Awesome-ML-SYS-Tutorial
My learning notes for ML SYS.
相关推荐
NeuPAN安装与配置:从零开始部署机器人导航框架的完整教程
NeuPAN安装与配置:从零开始部署机器人导航框架的完整教程 NeuPAN是一个基于端到端模型学习的直接点机器人导航框架,为机器人导航提供了创新的解决方案。本文
Dask Array 分块并行数组:从内部设计到 chunk 调优与实战
Dask Array 分块并行数组:从内部设计到 chunk 调优与实战 导读 Dask Array 是 Dask 项目中用于大规模数组计算的模块,它通过 分块
大数据数据分析任务调度ctf-wiki 堆利用系列:House Of Force 原理、chunk 尺寸计算与实战
ctf wiki 堆利用系列:House Of Force 原理、chunk 尺寸计算与实战 House Of Force(下称 HOF)是 glibc ptm
文档网络安全教程
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考