llama-trl
用 PPO 强化学习 + LoRA 高效微调 LLaMA 的三阶段 RLHF 训练框架
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
用 PPO 强化学习 + LoRA 高效微调 LLaMA 的三阶段 RLHF 训练框架
加载项目详情…
本应用为开源项目,仅供学习研究,请遵守其开源协议。
想象你训练一只猎犬:第一步教它基本口令(监督学习),第二步让它知道哪些动作会得到奖励(奖励模型),第三步让它自己探索最优策略(强化学习)。LLaMA-TRL 正是这样一套渐进式的微调方法论,它让 Meta 开源的 LLaMA 模型从"语言模仿者"进化为"有判断力的对话者"。
图1:LLaMA-TRL 三阶段训练流程
整个训练流程分为三个环环相扣的阶段:
在 supervised_finetuning.py 中,代码使用 HuggingFace TRL 库的 SFTTrainer 对 LLaMA 进行指令微调。核心配置:
from peft import LoraConfig
lora_config = LoraConfig(
r=8,
lora_alpha=16,
target_modules=["q_proj", "v_proj"],
lora_dropout=0.05,
bias="none",
task_type="CAUSAL_LM"
)
LoRA 的核心思想是"低秩分解":不直接更新原始大模型的权重矩阵 W,而是训练两个小矩阵 A 和 B,最终权重为 W + BA。由于 rank r=8,BA 的参数量仅为原始 W 的千分之一左右,却能保留 80% 以上的微调效果。target_modules 只指定 q_proj 和 v_proj(注意力机制的 Query 和 Value 投影),这正是 LoRA 的经典配置。
数据来源是 GPT-4-LLM 团队开源的 Alpaca 数据集(alpaca_gpt4_data.json),包含约 5.2 万条由 GPT-4 生成的指令-响应对。训练使用 cosine 学习率调度器,基础学习率 1e-5,最大步数 4000。
如果硬件充足,代码还提供了全参数微调版本 supervised_finetuning_full_weight.py,通过 DeepSpeed ZeRO Stage 3 实现了 CPU offload,可在有限 GPU 内存下完成全量权重更新。
training_reward_model.py 训练一个 Reward Model,本质是一个 AutoModelForSequenceClassification,输出一个标量分数来评价"回复质量"。这个模型的输入是一对「问题 + 回复」,输出是 [bad, good] 两个 logits,通过交叉熵损失训练——好的回复 logit 应当更高。
训练数据来自 comparison_data.json,包含多个候选回复的人工排序数据(如 GPT-4 生成的回复 vs. 人工优化的回复)。Reward Model 的作用至关重要:它决定了第三阶段 PPO 中"什么是好的回答"。
关键实现中,RewardTrainer 继承了 Trainer,使用了 HuggingFace 的 PeftModel 包装:冻结 base model,只训练 reward head + LoRA adapter,这大幅降低了训练成本。
tuning_lm_with_rl.py 是整个框架的核心创新所在。使用 TRL 库的 PPOTrainer 和 AutoModelForCausalLMWithValueHead:
from trl import PPOConfig, PPOTrainer, AutoModelForCausalLMWithValueHead
ppo_config = PPOConfig(
model_name=base_model_name,
learning_rate=1.4e-5,
mini_batch_size=1,
batch_size=8,
ppo_epochs=4,
gradient_accumulation_steps=4
)
PPO(Proximal Policy Optimization)的工作机制如下:
LoRA 在 PPO 阶段同样被使用,PPOTrainer 内部会调用 PeftModel 对 base model 进行包装,确保显存占用可控。
| 模块 | 功能 | 技术亮点 |
|---|---|---|
supervised_finetuning.py | SFT 监督微调 | TRL SFTTrainer + LoRA,8 GPU 并行 |
training_reward_model.py | Reward Model 训练 | PEFT + SequenceClassification |
tuning_lm_with_rl.py | PPO 强化学习微调 | TRL PPOTrainer + ValueHead |
utils/merge.py | LoRA 权重合并 | PeftModel.merge_and_unload() |
configs/default_offload_opt_param.json | DeepSpeed ZeRO-3 配置 | CPU offload 全量微调 |
技术栈覆盖:PyTorch、Transformers、PEFT、TRL、DeepSpeed、Accelerate、bitsandbytes(8-bit 量化)、WandB(训练监控)。
以 7B 参数 LLaMA 模型为例:
训练时间(7B 模型,8x A100):
LLaMA-TRL 填补了一个重要的技术空白:在 TRL 官方尚未原生支持 LLaMA 之前,它提供了完整的「从零用 RLHF 微调 LLaMA」的参考实现。2024 年初,HuggingFace TRL 正式合并了 PEFT 原生支持,使得 LoRA + PPO 成为主流范式,但 LLaMA-TRL 的三阶段设计(GPT-4 数据 → SFT → Reward Model → PPO)至今仍是 RLHF 论文和工业实践的基准模板。
项目在 GitHub 上获得了 239 颗星,被广泛引用于学术论文和开源教程。它证明了:
requirements.txt 使用 git+https://... 直接引用 GitHub 最新版,可能随时间出现依赖冲突