在 24GB 消费级 GPU 上对 20B LLM 进行 RLHF 微调
简要概述
Hugging Face 宣布 trl 库现在可以与 peft 和 8 位量化一起使用,允许在单个 24 GB 消费级 GPU 上对 20 B 参数的 LLM 进行 RLHF 微调。这使得大规模的 RL 微调变得经济且易于使用,无需多 GPU 模型并行设置。
为什么 RLHF 需要高效的微调
人类反馈强化学习(RLHF)通常包括三个阶段:(1) 基于指令的监督微调,(2) 从人工标注中训练奖励模型,(3) 使用奖励模型进行基于 PPO 的 RL 微调。RL 步骤需要在每个 GPU 上放置模型的两个副本(活跃模型和参考模型),这在模型参数超过 10 B 时很快就会超出单个设备的显存。
TRL:用于基于 PPO 的 RL 的库
trl 提供了用于语言模型 PPO 训练的高级 API。它利用 🤗 Accelerate 在单设备或分布式环境中运行。PPO 循环需要同时拥有活跃模型(正在更新)和参考模型(保持冻结),以计算 KL 正则化奖励,从而使显存占用翻倍。
使用 PEFT 和 8 位量化降低显存占用
8 位矩阵乘法
- 8 位量化(LLM.int8())将权重存储为每个参数 1 字节,相比 float32 将模型大小缩小四倍。
- 该方法将每个线性层拆分为处理异常值的 float16 部分和大部分的 int8 部分,在保持精度的同时提升速度。
通过 PEFT 的低秩适配(LoRA)
- LoRA 冻结预训练权重,并在注意力块的 query 和 value 投影中注入低秩矩阵(A 和 B)。
- 仅适配器参数可训练,从而显著降低优化器显存占用。
- 由于额外的矩阵乘法,前向和反向传播大约慢两倍,但显存节省使得在消费级硬件上训练 20 B 模型成为可能。
在 24 GB GPU 上对 20 B 模型的端到端流水线
步骤 1 – 以 8 位精度加载模型
model = AutoModelForCausalLM.from_pretrained(
"EleutherAI/gpt-neox-20b",
load_in_8bit=True,
device_map="auto",
)
以 8 位加载将显存需求从约 80 GB(float32)降低到约 20 GB,能够轻松适配 24 GB 显卡。
步骤 2 – 使用 PEFT 添加可训练的 LoRA 适配器
from peft import get_peft_model, LoraConfig
config = LoraConfig(r=8, lora_alpha=32, target_modules=["q_proj", "v_proj"], bias="none")
model = get_peft_model(model, config)
仅低秩矩阵会被存储在优化器状态中,将优化器显存占用从数 GB 降至数百 MB。
步骤 3 – 使用单一模型生成参考和活跃 logits
PEFT 的 disable_adapters 上下文管理器会暂时停用 LoRA 层,使同一底层模型能够生成参考 logits:
with model.disable_adapter():
ref_logits = model(input_ids)
# active logits are computed with adapters enabled
active_logits = model(input_ids)
无需第二个完整模型副本,进一步降低显存使用。
训练脚本概览
博客文章提供了三个脚本,演示在 20 B GPT‑NeoX 模型上的完整工作流:
clm_finetune_peft_imdb.py– 在 IMDB 情感数据集上对 LoRA 适配器进行因果语言模型微调(一个 epoch)。merge_peft_adapter.py– 将 LoRA 权重合并到基础模型中,以用于推理或进一步训练。gpt-neo-20b_sentiment_peft.py– 使用 IMDB 情感分类器作为奖励模型进行 PPO 微调,以生成正面的电影评论。
所有脚本均在 NVIDIA RTX 4090(24 GB)上执行。完整的训练也在 🤗 研究集群中的单个 A100 上进行测试。
结果
- 损失曲线显示在 IMDB 上进行一次 epoch 的监督 LoRA 微调后实现了稳定收敛。
- 在 PPO 过程中,平均奖励持续上升,表明模型学会生成更积极的评论。
- 整个流水线在单个 24 GB GPU 上运行,证明 RLHF 已不再局限于多 GPU 集群。
对社区的意义
- 降低入门门槛 – 研究人员和开发者可以在消费级硬件上尝试 RLHF。
- 开源可复现性 – 所有代码和适配器均托管在 Hugging Face Hub,便于共享微调产物。
- 可扩展的基础 – 一旦加入多 GPU 支持,同样的方法可通过数据并行扩展到更大的模型。
未解问题与未来工作
- 多 GPU 扩展 – 该集成在跨多 GPU 的数据并行下表现如何?
- 训练速度 – LoRA 带来额外开销;探索更快的 kernel 或混合精度策略可能缓解此问题。
- 更广泛的 RL 算法 – 虽然 PPO 是默认选择,集成其他 RL 方法(如 DPO)可能扩大适用范围。
参考文献
- 并行范式 – https://huggingface.co/docs/transformers/v4.17.0/en/parallelism
transformers中的 8 位集成 – https://huggingface.co/blog/hf-bitsandbytes-integration- LLM.int8() 论文 – https://arxiv.org/abs/2208.07339
- 梯度检查点 – https://docs.aws.amazon.com/sagemaker/latest/dg/model-parallel-extended-features-pytorch-activation-checkpointing.html