9.0
重磅
AI SCORE
技术实践2024-02-11 23:12
50行Python代码实现LLM的RLHF微调
Hacker News · datadreamer.dev#RLHF#微调#教程
Editor brief · 编辑速览
用极简代码详解RLHF训练完整流程。对想快速理解强化学习微调原理或快速原型验证的程序员极其实用。
为了更好地将指令微调的 LLM 生成的回应与人类的偏好相对齐,我们可以针对奖励模型或人类偏好数据集来训练 LLM,这个过程被称为 RLHF(带人工反馈的强化学习)。
DataDreamer 使这个过程变得极其简单直接。我们在下面展示了一个示例,使用 LoRA 只训练一部分权重,并采用 DPO(一种比传统 RLHF 更稳定、更高效的对齐方法)。
from datadreamer import DataDreamer
from datadreamer.steps import HFHubDataSource
from datadreamer.trainers import TrainHFDPO
from peft import LoraConfig
with DataDreamer("./output"):
# Get the DPO dataset
dpo_dataset = HFHubDataSource(
"Get DPO Dataset", "Intel/orca_dpo_pairs", split="train"
)
# Keep only 1000 examples as a quick demo
dpo_dataset = dpo_dataset.take(1000)
# Create training data splits
splits = dpo_dataset.splits(train_size=0.90, validation_size=0.10)
# Align the TinyLlama chat model with human preferences
trainer = TrainHFDPO(
"Align TinyLlama-Chat",
model_name="TinyLlama/TinyLlama-1.1B-Chat-v1.0",
peft_config=LoraConfig(),
device=["cuda:0", "cuda:1"],
dtype="bfloat16",
)
trainer.train(
train_prompts=splits["train"].output["question"],
train_chosen=splits["train"].output["chosen"],
train_rejected=splits["train"].output["rejected"],
validation_prompts=splits["validation"].output["question"],
validation_chosen=splits["validation"].output["chosen"],
validation_rejected=splits["validation"].output["rejected"],
epochs=3,
batch_size=1,
gradient_accumulation_steps=32,
)