ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

slime PPO 性能优化实战:GAE 分 chunk 并行计算(Chunk-Scan)原理、推导与工程落地

slime PPO 性能优化实战:GAE 分 chunk 并行计算(Chunk-Scan)原理、推导与工程落地 文档教程人工智能大模型RLHF【免费下载链接】Awesome-ML-SYS-TutorialMy learning notes for ML SYS.项目地址https://gitcode.com/gh_mirrors/aw/Awesome-ML-SYS-Tutorial点击查看免费下载本篇技术指南围绕开源仓库 rlhf/slime/batch-GAE/ppo-gae-chunk.md 所记录的 slime 框架 PPO 训练管线改造展开面对 agentic RL 长序列场景下 GAEGeneralized 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 计算时间得到100×–300×的加速实测最大约317×并行度取决于chunk_size在不 OOM 的前提下chunk_size越大加速越明显。需要说明本文依据仓库中的学习笔记整理涉及的外部技术报告与上游 PR 细节以仓库记录为准。slime 框架本身是基于 SGLang 与 Megatron LM 作为唯一后端的 RL 训练框架其整体架构与数据流可参考仓库中 rlhf/slime/code-walk-through/readme.md 的代码走读记录。2. 技术背景为什么要搞 GAE 的 Chunk-Scan2.1 为什么现在 GAE 会变成瓶颈在 RLHF / Agentic RL 里PPO 仍是一个非常常用、表现稳定的算法。我们需要在每个 token 上计算 advantage最常见的就是 GAE。而 GAE 的标准写法是一个从后往前的递推公式对序列长度 T 来说是 O(T) 的串行依赖——A_t依赖A_{t1}无法直接并行。在 slime 中GAE 的算法实现如下lastgaelam torch.zeros(B, devicedevice, dtypedtype) 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], dim1) # [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_{kt}^{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 GAE3.1 标准 GAE 回顾在开始前先回顾一下标准的 GAE。记 delta 为$$ \delta_t r_t \gamma V_{t1} - V_t $$则 GAE 的 advantage 为$$ A_t \sum_{kt}^{T-1} (\gamma \lambda)^{k-t} \delta_k $$也可以写成后向递推的形式$$ A_t \delta_t \gamma \lambda A_{t1}, \quad t T-1, T-2, \dots, 0 $$可以看到后向递推形式天然是串行的每个时间步都要等它后面一个时间步的结果。三种解法的差异本质上就是如何打破这条串行依赖链。3.2 方案一串行解法slime 目前的版本给出的答案与 2.1 节同一段代码lastgaelam torch.zeros(B, devicedevice, dtypedtype) 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], dim1) # [B, max_len]优点实现简单数值稳定缺点这个版本在时间维度完全串行长序列下性能不行。值得注意的细节是next_value的边界处理当t max_len - 1最后一个 token时没有V_{t1}按 0.0 处理adv_rev按逆序收集结果后通过torch.stack(adv_rev[::-1], dim1)再翻转为正向时间顺序得到形状为[B, max_len]的完整 advantage 矩阵。这个边界约定在后续所有并行化方案中都必须保持一致。3.3 方案二纯矩阵解法利用前向展开式$$ A_t \sum_{kt}^{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第一个 chunk0 - C-1 第二个 chunkC - 2C-1 ... 第 c 个 chunkcC - (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}\cdotsw^t \widetilde{\delta}_{cC}w^{t1} S_{cC - 1} \end{aligned} $$把当前 chunk 内的部分单独拿出来$$ s^{(c)}t \sum{k0}^{t} w^{t-k} \widetilde{\delta}_{cC k} $$于是最终公式可以写成$$ \boxed{ S_{cC t} s^{(c)}t w^{t1} , s{\text{prev}}, \quad t 0, \dots, L_c - 1 } $$这意味着局部部分$s^{(c)}_t$ 可以在 chunk 内用矩阵/conv 并行算——它只依赖当前 chunk 内部的 $\widetilde{\delta}$是 chunk 内部的前缀加权和形式上与线性 attention 的 chunk 内扫描完全同构跨 chunk只需要维护一个标量状态s_prevchunk 之间串行递推即可每个 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_{t1} - 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^(t1) * s_prev S_global s_local s_prev * pow_vec[0:L_c] write_into(full_scan_rev, chunk_indexc, valuesS_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对这段伪代码的几个实现要点展开说明kernelC×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^{t1}$。注意 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 记录了 slimeTraining (Megatron) Rollout (SGLang) Data Buffer的分离式架构可以帮助理解 GAE 计算在整个流水线中的位置。4. 实现效果根据仓库笔记记录的实验结果实现效果非常可观No chunkchunk size 64chunk size 128chunk size 256B256, T1310725.935994s0.070122s0.034059s0.018390s ( x317 )B128, T655362.902570s0.232986s0.017645s0.009134s可以看到在T131072128K 超长序列、chunk_size256时加速比约317×——从接近 6 秒压到不足 20 毫秒在T65536、chunk_size256时加速比同样非常可观——从 2.9 秒压到约 9 毫秒。结论只要有足够显存来提升 chunk size并行度就能大幅增加GAE 的计算时间也能被相当可观地缩减。这里也解释了 chunk_size 的调参方向chunk_size 越大chunk 数越少chunk 间的串行递推步数越少、并行度越高但 chunk 内的 C×C 核矩阵空间复杂度 O(C²)占用显存也越大因此需要在显存预算内尽量增大 chunk_size。作为参考表中chunk_size64与chunk_size128、chunk_size256的耗时递减趋势与这一规律一致。5. 具体使用方法在 slime 里怎么用 Chunk-Scan GAEChunk-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²) 虽然完美并行却必然 OOMChunk-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-TutorialMy learning notes for ML SYS.项目地址https://gitcode.com/gh_mirrors/aw/Awesome-ML-SYS-Tutorial点击查看免费下载相关推荐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),仅供参考
返回列表