RLHF对齐训练的基本流程
RLHF(Reinforcement Learning from Human Feedback)是大语言模型对齐训练的核心方法,通过人类偏好数据引导模型输出符合预期的回答。整个人工智能对齐训练流程分为三个阶段:监督微调(SFT)、奖励模型训练(RM)、强化学习优化(PPO)。监督微调阶段使用高质量对话数据对预训练模型进行指令微调,奖励模型阶段训练一个能够预测人类偏好的打分模型,PPO阶段则利用奖励模型的反馈信号优化策略模型。
大模型开发中,RLHF的关键挑战在于奖励模型的泛化能力和PPO训练的稳定性。奖励模型如果过拟合训练分布,会导致策略模型在未见过的输入上产生奖励黑客(reward hacking)现象——模型学会钻奖励模型的漏洞而非真正提升回答质量。
奖励模型训练方法与数据构造
奖励模型的输入是(prompt, response)对,输出是一个标量奖励值。训练数据通常通过人工标注获得:对同一个prompt生成多个回答,标注员对回答进行排序,排序差异转化为Bradley-Terry偏好对。
奖励模型的损失函数基于偏好对构造:
import torch
import torch.nn as nn
class RewardModel(nn.Module):
def __init__(self, backbone_model):
super().__init__()
self.backbone = backbone_model
self.reward_head = nn.Linear(backbone.config.hidden_size, 1)
def forward(self, input_ids, attention_mask):
outputs = self.backbone(
input_ids=input_ids,
attention_mask=attention_mask
)
last_hidden = outputs.last_hidden_state[:, -1, :]
reward = self.reward_head(last_hidden)
return reward.squeeze(-1)
def preference_loss(reward_chosen, reward_rejected):
# Bradley-Terry偏好损失
logits = reward_chosen - reward_rejected
loss = -torch.nn.functional.logsigmoid(logits).mean()
return loss
def train_reward_model(model, dataloader, optimizer, epochs=3):
model.train()
for epoch in range(epochs):
total_loss = 0
for batch in dataloader:
chosen_ids = batch["chosen_input_ids"]
chosen_mask = batch["chosen_attention_mask"]
rejected_ids = batch["rejected_input_ids"]
rejected_mask = batch["rejected_attention_mask"]
reward_chosen = model(chosen_ids, chosen_mask)
reward_rejected = model(rejected_ids, rejected_mask)
loss = preference_loss(reward_chosen, reward_rejected)
optimizer.zero_grad()
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
optimizer.step()
total_loss += loss.item()
print(f"Epoch {epoch+1}, Loss: {total_loss/len(dataloader):.4f}")
奖励模型的backbone通常使用与策略模型相同架构的预训练模型,去掉LM head后接一个线性输出层。训练数据量一般需要5万到20万条偏好对才能获得较好的泛化效果。
PPO算法核心实现
PPO(Proximal Policy Optimization)是RLHF中最常用的强化学习算法。与标准RL不同,RLHF中的PPO需要同时维护四个模型:策略模型(Actor)、参考模型(Reference,冻结的SFT模型)、奖励模型(冻结的RM)、critic模型(Value Model)。参考模型用于计算KL散度惩罚,防止策略模型偏离SFT模型太远。
import torch
import torch.nn.functional as F
class PPOTrainer:
def __init__(self, actor_model, critic_model, reference_model,
reward_model, clip_ratio=0.2, kl_coef=0.2):
self.actor = actor_model
self.critic = critic_model
self.reference = reference_model
self.reward_model = reward_model
self.clip_ratio = clip_ratio
self.kl_coef = kl_coef
for p in self.reference.parameters():
p.requires_grad = False
for p in self.reward_model.parameters():
p.requires_grad = False
def compute_advantages(self, rewards, values, gamma=0.99, lam=0.95):
# GAE广义优势估计
advantages = []
gae = 0
for t in reversed(range(len(rewards))):
if t == len(rewards) - 1:
next_value = 0
else:
next_value = values[t + 1]
delta = rewards[t] + gamma * next_value - values[t]
gae = delta + gamma * lam * gae
advantages.insert(0, gae)
advantages = torch.tensor(advantages)
advantages = (advantages - advantages.mean()) / (advantages.std() + 1e-8)
return advantages
def train_step(self, prompts, responses, old_log_probs):
with torch.no_grad():
rewards = self.reward_model(prompts, responses)
ref_log_probs = self._compute_log_probs(
self.reference, prompts, responses)
cur_log_probs = self._compute_log_probs(
self.actor, prompts, responses)
kl_penalty = (cur_log_probs - ref_log_probs).mean(dim=-1)
rewards = rewards - self.kl_coef * kl_penalty
values = self.critic(prompts, responses)
advantages = self.compute_advantages(
rewards.tolist(), values.tolist())
new_log_probs = self._compute_log_probs(
self.actor, prompts, responses)
ratio = torch.exp(new_log_probs - old_log_probs)
clipped_ratio = torch.clamp(
ratio, 1 - self.clip_ratio, 1 + self.clip_ratio)
policy_loss = -torch.min(
ratio * advantages,
clipped_ratio * advantages
).mean()
value_loss = F.mse_loss(values, rewards.detach())
loss = policy_loss + 0.5 * value_loss
return loss, policy_loss.item(), value_loss.item()
def _compute_log_probs(self, model, prompts, responses):
outputs = model(input_ids=responses, attention_mask=None)
logits = outputs.logits[:, :-1, :]
labels = responses[:, 1:]
log_probs = F.log_softmax(logits, dim=-1)
token_log_probs = log_probs.gather(
-1, labels.unsqueeze(-1)).squeeze(-1)
return token_log_probs
训练参数调优与常见问题
RLHF训练中最关键的参数是KL系数(kl_coef)和PPO clip ratio。KL系数控制策略模型偏离参考模型的程度——值太大模型几乎不更新,值太小模型容易崩溃。实践中通常从0.2开始,配合自适应KL调整策略:当KL散度超过目标值时增大系数,低于目标值时减小系数。
clip ratio一般设为0.2,这是PPO论文中的默认值。如果训练不稳定,可以降低到0.1。学习率方面,策略模型通常用1e-6到5e-6,critic模型用1e-5到5e-5,两者使用不同的学习率是因为critic需要更快收敛来提供准确的价值估计。
奖励黑客是RLHF中最棘手的问题。检测方法是在验证集上跟踪奖励模型的分数与人类真实偏好的相关性变化趋势。如果奖励分数持续上升但人类评估分数下降,说明出现了奖励黑客。解决方案包括增大KL惩罚系数、增加奖励模型的正则化、引入多样性奖励等。
显存方面,四个模型同时加载是RLHF的主要瓶颈。对于7B级别的模型,全精度训练需要约200GB显存。使用LoRA微调可以将显存降至约80GB,配合DeepSpeed ZeRO-3或FSDP分片策略可以进一步降低到单机8卡可训练的水平。
原创文章,作者:小编,如若转载,请注明出处:https://www.yunthe.com/rlhf-qiang-hua-xue-xi-dui-qi-xun-lian-shi-zhan-ppo-suan-fa/