强化学习框架实战:基于 Miles 优化大语言模型训练
强化学习框架实战:基于 Miles 优化大语言模型训练
大模型 SFT 之后,怎么把 RLHF 跑稳?这是兰屿星奇最近啃的一块硬骨头。我们内部研发的 Miles 框架(注:非开源,基于 slime 原型迭代而来)在这次项目中帮了大忙。它重写了通信层和内存分配器,处理多卡同步时没那么容易崩。下面聊聊怎么用它跑通 PPO,以及踩过的坑。
Miles 是什么
Miles 不是开源项目,是我们基于早期 slime 实验代码重构的内部工具。主要解决两个问题:大显存占用,以及分布式训练里的通信瓶颈。
架构上沿用 Actor-Critic 分离模式,但底层换成了更轻量的通信原语。开发者不用手撸 NCCL,只要写好奖励逻辑就行。这点对于只有少量 GPU 资源的团队很友好。
训练流程配置
代码层面,Miles 的接口还算直白。下面这段是我们在 8×A100 集群上实际跑通的基线配置,单 epoch 耗时约 11 小时:
from miles import Trainer, Policy, RewardModel
from transformers import AutoModelForCausalLM
# 1. 加载 SFT 后的策略模型
policy = Policy.from_pretrained("your-sft-model-path")
# 2. 加载奖励模型
reward_model = RewardModel.from_pretrained("your-rm-model-path")
# 3. 配置训练环境
trainer = Trainer(
policy=policy,
reward_model=reward_model,
learning_rate=1e-5,
kl_coefficient=0.1, # 防止策略偏离 SFT 模型过远
batch_size=32,
gradient_accumulation_steps=4
)
# 4. 开始训练
trainer.train(dataset=your_preference_dataset, epochs=3)
关键在 kl_coefficient。Miles 内部会对这个值做自动缩放,适配不同的序列长度,不用手动调参。
常见陷阱与性能调优
OOM 了。
长上下文场景下,Attention 机制依然是显存杀手。解决办法很简单:开启 flash_attention_2,然后把 max_seq_len 卡死,别让没必要的 padding 撑爆显存。
奖励信号稀疏。
奖励模型如果反馈太粗,策略更新就会瞎走。我们的做法是先上 Rejection Sampling 筛一遍优质样本,再丢进 Miles 跑 PPO。这样收敛快,也更稳。
别信默认值。超参数必须根据自家数据重新调,尤其是 learning_rate 和 batch size。
写在最后
文档里那个 custom_loss.py 示例,我调了三天才跑通——欢迎来钉钉骂我。
本文首发于 强化学习框架实战:基于 Miles 优化大语言模型训练 — https://lyxq.com.cn/en/blog/miles-rl-framework-optimization
转载或引用请注明出处,商业使用请联系作者获得授权。