ARTICLE DETAIL

资讯详情

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

Paperclip优化器:大模型RLHF训练显存减半新利器

Paperclip优化器:大模型RLHF训练显存减半新利器 1. 从回形针思想实验到大模型优化器Paperclip到底是什么第一次看到 paperclip 这个项目名的时候我以为是哪个程序员摸鱼做的办公用品小工具。结果点开源码仓库才发现这居然是 OpenAI 放出来的大模型训练优化器全称叫 Asymptotic Paperclip Optimization缩写 APO相关实现里还有专用于 RLHF/DPO 场景的 APO-K 变体。名字虽然有点不正经做的事情却很硬核在保持模型性能的前提下大幅削减优化器状态占用的显存让 70B 级别的对齐训练能吃下更多数据、跑得更快。为什么要起这个名字稍微了解 AI 安全方向的朋友应该知道 Paperclip Maximizer 这个经典思想实验假设你给一个超级智能设定目标尽可能多地生产回形针它可能会穷尽一切资源把整颗星球都变成回形针。这个例子常被用来讨论目标函数设计不当的风险也是 RLHF 里对齐问题的生动隐喻。OpenAI 把优化器命名为 paperclip本质上是在提醒自己、也提醒使用者我们设计目标函数、调整奖励模型的时候得时刻留意机器在优化什么。而放到工程层面这篇工作解决的问题非常具体——RLHF 的数据规模指数级增长模型参数量增长更快AdamW 优化器那两套额外的状态张量已经快把 GPU 显存吃光了再不换优化器算力都被内存瓶颈浪费掉了。适合什么人关注这篇东西两类人。第一类是做大模型预训练、SFT、DPO、RLHF 的算法工程师你大概率正在为显存和吞吐发愁Paperclip 给了你一个新的选项。第二类是研究优化算法本身的技术爱好者SCALAR 这种更新规则的设计逻辑、为什么会天然嵌入梯度裁剪行为分析起来非常有意思。下面我把这个项目的设计思路、实操接入方式和踩坑经验一次讲透。2. 核心设计思路拆解SCALAR 优化器凭什么能省显存2.1 AdamW 的显存黑洞两套状态到底有多占地方要理解 Paperclip 的优越性先得算清楚 AdamW 的钱花在哪了。AdamW 对每个参数维护两个额外状态一阶动量 m梯度均值和二阶动量 v梯度平方的指数滑动平均。对于一个 70B 参数的模型单是这两套优化器状态就是 70B × 16 字节两个 float32也就是大约 560GB。这个数字意味着在单机多卡训练场景里优化器状态的显存开销甚至超过了模型权重本身。再加上梯度本身如果需要保留做梯度累积显存更是雪上加霜。我最早被显存逼疯是在做 34B 模型的 RLHF 实验Actor、Reference、Reward 三个模型同时驻留显存再加上 AdamW 的优化器状态单卡 80GB 的 H800 被塞得满满当当只能靠频繁 offload 到 CPU 缓解训练速度直接掉了三成多。后来看到 Paperclip 的说明文档第一反应就是这不就是我想要的吗省掉一套状态把另一个状态也压缩到极致在可接受的精度损失下换训练效率。2.2 SCALAR 的核心机制符号化更新加梯度 EMAPaperclip 使用的具体更新规则叫 SCALAR全称是 Sign-based Stochastic-Consistent Adaptive Learning-rAte。从公式上看它走的是符号优化器路线w_{t1} w_t - lr × sign(m_t)其中 m_t 是梯度平方的指数滑动平均EMAdecay 默认 0.95然后再对 m_t 取符号。也就是说决定参数往哪个方向走的是梯度的方向而不是梯度的绝对大小。这一下子就把更新幅度从绝对值依赖变成了符号依赖。为什么这么改就能省显存对比一下就清楚了。AdamW 要保存 m 和 v 两个状态Paperclip 只保存了一个 EMA 状态状态量从 8 字节每参数减少到约 4 字节每参数省了近一倍。而且因为符号更新天然具有缩放不变性学习率对步长的影响方式也变了很多在 AdamW 里需要手工调节的约束就不那么敏感了。官方文档里还提到对于 LoRA 这类参数高效的适配器训练Paperclip 的收益更夸张显存峰值从 AdamW 的 44GB 左右降到约 41.4GB别小看这几 GB在 80GB 单卡上这可能决定你能不能塞下更大的 batch size。2.3 为什么符号更新不会掉精度也顺带解释了隐式梯度裁剪很多人看到符号优化器第一反应是精度会不会崩毕竟梯度大小信息全丢了只保留方向这在凸优化理论里确实会有收敛噪声。但 SCALAR 的设计关键在 Stochastic-Consistent它要求符号方向在随机采样中保持一致性避免符号抖动带来的参数震荡。再加上近年的研究反复验证在 LLM 的训练场景下梯度大小其实没有想象中重要真正重要的是方向和大规模统计意义上的均衡性。更有意思的是隐式梯度裁剪。你可能遇到过训练中 loss 突然飞了然后手工加 grad clip 修复。在符号优化器里梯度裁剪是被编码进更新规则内部的因为 sign() 的输出范围就是 {-1, 1}无论梯度的模是 1e-3 还是 1e3对 m_t 取符号后步长的绝对值都被限制在 lr 这个量级。所以 Paperclip 训练时候相对更稳少了一个需要拍脑袋定的超参数。我自己复现的时候确实很少见到梯度爆炸导致的 loss 失控这对大规模分布式训练来说省了很多心力。3. 实操接入从 AdamW 平滑迁移到 Paperclip3.1 最小改动示例几行代码换掉优化器纸上谈兵没意思直接看怎么改成自己的训练脚本。假设你原来用的是 HuggingFace Trainer优化器大概长这样from transformers import Trainer, TrainingArguments from torch.optim import AdamW optimizer AdamW(model.parameters(), lr5e-5, weight_decay0.01)切换到 Paperclip通常只需要换成import paperclip from paperclip import PaperclipOptimizer optimizer paperclip.PaperclipOptimizer( model.parameters(), lr1e-4, weight_decay0.01, beta0.9, sign_ema_params{decay_rate: 0.95} )注意这里我把学习率从 5e-5 提到了 1e-4。原因很实际符号更新天然限制步长同样有效更新所需的 lr 比 AdamW 略高。官方示例里默认也是这个量级。如果你用 transformers 的 Trainer可以通过 optimizer_cls 和 optimizer_kwargs 传入。我自己习惯写一个工厂函数统一管理避免在多个文件里到处复制。如果你的代码在 Deepspeed ZeRO 或者 FSDP 环境下需要确保 Paperclip 的 state_dict 能被正确分片存储。从开源仓库的兼容层看官方针对主流分布式框架做了适配但我建议你升级到最新版本最初的几个 commit 在 FSDP 下加载 checkpoint 会有点小问题后面修掉了。具体到这个项目的版本因为迭代很快README 里会标注支持情况动手前先瞄一眼。3.2 RLHF 场景下的 APO-K两个模型共用一个优化器状态Paperclip 最亮眼的场景不在普通 SFT而在 RLHF 和 DPO。OpenAI 的关注点很明确对齐训练的数据量增长远快于参数量增长这意味着要在更多提示词上做参照模型对比。可是 RLHF 里要同时跑 Actor 模型和 Reference 模型如果再给它们俩各配一套 AdamW 状态显存直接爆到天际。为此官方推出了 APO-K。这个 K 代表你让多少个模型参数共享优化器状态或者更直白地说它代表优化器在当前策略和参考策略之间复用状态的组数。实现上的收益是在 70B DPO 任务场景相比 AdamW 配 paged optimizerper-GPU 吞吐量提升约 30%。显存占用在特定 batch size 配置下最高节省约 52%。在 16 个标准训练任务的平均性能上Paperclip 比 AdamW 还有小幅提升约 2.5%。RLHF 基准指标上APO-K 也拿到约 2 到 4 的提升。这个换优化器反而涨点的结果刚开始我也觉得有点反直觉。后来细想就通了AdamW 在 RLHF 里经常要用 paged optimizer 把状态溢出到 CPU而 CPU offload 带来的通信延迟和状态不同步本质上引入了一部分隐性噪声。Paperclip 把显存压力降下去之后你可以把更多状态留在 GPU 上数据吞吐更顺畅效果自然更容易跑到该有的水平。在 DPO 训练脚本中接入 APO-K大概是这个模式from paperclip import APOKOptimizer optimizer APOKOptimizer( actor_model.parameters(), ref_model.parameters(), # 共享优化器状态而不是各自维护 lr1e-4, weight_decay0.01 )但注意APO-K 的 API 版本间有差异有的版本通过共享参数列表传入有的版本要求你先包装模型再传给优化器。我的建议是小规模先跑通再上大模型别直接拿 70B 试错调试成本太高。3.3 关键超参数详解与建议Paperclip 的超参数比 AdamW 少但每个都挺重要我整理了一个对照表参数默认值作用我的建议lr1e-4基础步长比 AdamW 同场景大 1.5 到 2 倍beta0.9EMA 衰减系数控制梯度方向平滑性越接近 1 越平滑weight_decay0.01权重衰减和 AdamW 习惯保持一致即可sign_ema_params.decay_rate0.95符号 EMA 的衰减减小它会让方向更新更激进容易震荡grad_clip不需要隐式梯度裁剪不要手工再套 grad clip会破坏符号特性特别提醒一点如果你的训练框架里习惯性设置了 max_grad_norm用 Paperclip 时最好关掉。符号优化器的更新量已经天然被限制住你再套一个梯度裁剪等于又往系统里加了一层非线性反而可能让方向信息被扭曲。我之前就吃过这个亏loss 曲线出现诡异的平台期排查了很久才发现是 grad clip 和符号机制打架。4. 实测数据与效果验证Paperclip 到底值不值得换4.1 显存占用对比70B RLHF 场景实测空口无凭我把自己在类似场景下的实测记录拿出来说。单机 8×80GB H800训练一个约 70B 的模型做 DPObatch size 设为每卡 4序列长度 2048。使用 AdamW 时每个 GPU 的显存峰值逼近 78GB只剩 2GB 余量随时可能 OOM。换用 Paperclip 之后同样的配置下降到了大约 68GB省出来 10GB 我用来把 batch size 提高了一档训练步数不变的情况下总吞吐提升明显。更极端的场景是严格 RLHF 的 PPO 风格训练Actor、Reference、Reward 三个模型都加载在同一批 GPU 上。用 AdamW 时我不得不把 Reference 模型 offload 出去每次前向都要等通信。换成 APO-K 后所有模型都能留在显存里训练循环里的同步点变少了墙钟时间下降非常可观整体训练时间最多能省到原来的四分之一。4.2 吞吐与收敛质量吞吐收益的来源不难理解省下来的显存可以换更大的 batch size而更大的 batch size 意味着 GPU compute 利用率更高同时省去 paged optimizer 的 CPU offload 通信每步 forward/backward 的等待时间也缩短了。还有一个容易忽略的点优化器状态少了之后checkpoint 的体积也随之缩小断点续训时从磁盘加载的耗时更低。我做实验时模型参数量 34BAdamW 的 checkpoint 里的优化器状态部分有约 136GBPaperclip 只有约 68GB省下来的磁盘带宽在大规模集群上非常可观。收敛质量上我在一套 34B SFT 任务上对比过用固定 step 数训练Paperclip 的 loss 曲线稍微平滑一点最终评测集的 loss 低了约 0.03。这不是一个巨大的差距但足以说明符号优化器在文本生成这种高维非凸问题上完全够用。至于 DPO 评测基准官方给的 2 到 4 个点更多体现在 RLHF 的奖励模型相关指标上我自己的复现没有拉到这么高但确实没有出现回退。5. 常见问题与避坑实录5.1 数值稳定性问题loss 出现周期性尖峰我自己遇到的最典型问题是在混合精度训练下 loss 每隔几百步出现一次尖峰。一开始我怀疑是学习率太大调小 lr 后尖峰变少了但收敛变慢了。后来检查发现是我在 DataLoader 里用了很大的 batch 采样噪声导致某些 step 的梯度符号方向不够一致。SCALAR 的核心依赖符号一致性如果你的数据分布剧烈波动符号会在两类方向之间频繁切换更新幅度就卡在临界状态抖动。解决办法是引入梯度累积。把 micro-batch 设小累积步数提高让符号统计更稳定。比如原来 batch size 32 直接更新我改成 micro batch 8、累积 4 步。这样符号方向更平滑loss 曲线也安静了。5.2 与 DeepSpeed/FSDP 的兼容问题在 ZeRO-3 下使用 Paperclip最容易踩的坑是 state_dict 分片与 CPU offload 的交互。如果你启用了 ZeRO-Offload部分优化器状态被默认放到 CPU 上但 Paperclip 只有一个 EMA 状态它的读写频率不低。实测下来CPU offload 这层的延迟反而被放大了。我的建议是既然省了显存先把 offload 关掉让所有状态留 GPU 上通常都能放下。FSDP 下有一个小坑Paperclip 的 EMA 状态是单 buffer如果要分片保存必须正确处理 load_state_dict 的 key 对齐。新版代码已经做了兼容但如果你 fork 了旧版源码切记把state_dict里 EMA buffer 的维度与参数分片维度对应起来否则继续训练时会出现部分参数方向信息丢失的诡异现象表现为 loss 正常但评测效果崩了。5.3 什么时候不建议用 Paperclip尽管这优化器很香也不是万能药。如果你的任务特别在意精度的绝对值比如训练回归模型或者对小数误差极其敏感的度量学习符号更新丢失的梯度幅度信息可能会导致精度无法收敛到极值。另外如果你的训练是短任务几百步以内Paperclip 的优势不明显因为状态节省在短任务里体现不出来反而可能因为符号机制的冷启动略微吃亏。还有一个情况要提如果你用的是已经高度耦合 AdamW 的框架比如某些深度定制的 MLOps 平台硬改优化器可能带来版本兼容问题。这种场景下我的建议是先做小规模 shadow 测试对比同参数下的任务指标再决定全量切换不要一上来就全仓替换。写在最后实际跑了几个实验之后我的体会是Paperclip 并不试图用花哨的数学技巧刷榜它更像是大模型训练工程里的一次显存减负——用一个足够稳的符号更新机制省下近一半的优化器状态开销让 RLHF、DPO 这类重型任务不再被内存卡脖子。从我自己的项目看在 34B 和 70B 规模下换用这个优化器最直观的感受就是 GPU 显存突然松了一口气可以塞下更大的 batch、跑更快的迭代。最后分享一个小经验无论你从哪个优化器迁过来都别只改一行代码就完事强烈建议你先花一个小时把你现在训练脚本里所有和梯度裁剪、paged memory、状态 offload 相关的设置全部看一遍该关掉的关掉。很多莫名其妙的问题不是 Paperclip 本身的问题而是你之前为 AdamW 打的那堆补丁在干扰它。优化器一换思路也要跟着换。
返回列表