在大语言模型(LLM)的后训练阶段,基于人类反馈的强化学习(RLHF) 已成为提升模型对话能力和对齐性的关键步骤。其中,近端策略优化(PPO) 因其稳定性和效率被广泛采用。然而,将 PPO 扩展到百亿甚至千亿参数规模的模型时,显存、通信和训练稳定性都面临巨大挑战。
微软 DeepSpeed 团队在 DeepSpeedChat 中提供了 DeepSpeedPPOTrainer 类,它深度融合了 ZeRO-3 显存优化、混合精度训练 、梯度溢出处理 以及高效的经验生成与管理机制,为大规模 RLHF 提供了工业级的解决方案。本文将从代码结构出发,逐层剖析该训练器的设计思想与实现细节。
一、整体架构与模块职责
DeepSpeedPPOTrainer 是一个集成了四个子模型的协调器:
- Actor 模型:待优化的策略模型(即我们最终要输出的对话模型)。
- Reference 模型:Actor 的冻结副本,用于计算 KL 散度惩罚,防止策略偏离太远。
- Critic 模型:价值网络,估计状态值函数,用于计算优势函数(GAE)。
- Reward 模型:预先训练好的奖励模型,为生成的回答打分。
初始化时,这些模型均通过 RLHFEngine 传入,并已配置好 DeepSpeed 引擎(含 ZeRO 策略、优化器等)。DeepSpeedPPOTrainer 本身不持有优化器,而是通过调用各个模型的 backward() 和 step() 方法驱动参数更新,因此天然支持多模型独立/联合训练。
核心超参数包括:
kl_ctl:KL 惩罚系数(默认 0.1)cliprange/cliprange_value:策略和价值网络的截断范围(均默认 0.2)gamma/lam:GAE 的折扣因子和 λ 参数clip_reward_value:奖励裁剪阈值,防止极端奖励值导致训练不稳定
二、经验生成:兼顾效率与鲁棒性
RLHF 的每一次迭代都需要从当前 Actor 采样一批回答,并计算相应的对数概率、参考概率、奖励和状态值。generate_experience 方法封装了这一完整流程。
1. 序列生成
_generate_sequence 调用 HuggingFace 的 generate 接口,并特别处理了 LLaMA 模型的采样问题(早期版本 do_sample 可能导致 NaN)。生成时通过 synced_gpus 参数支持 ZeRO-3 下的同步生成,避免跨 GPU 参数不一致。
2. 过滤无效回答
由于预训练模型可能产生过短或无意义的回答,代码会检查生成序列的有效长度(valid_ans_len <= 1 则丢弃)。若整批全部无效,则会复用 last_generated_experience ------ 这是一种优雅降级策略,防止因个别批次生成失败导致训练中断。
3. 一次性前向计算
生成后,利用 Actor、Reference、Reward 和 Critic 四个模型并行前向,分别得到:
logprobs和ref_logprobs(通过gather_log_probs从 logits 中提取)reward_score(来自 Reward 模型的chosen_end_scores)values(Critic 输出,去掉最后一个 token 的预测)
所有计算均在 torch.no_grad() 下进行,且返回值会被缓存,用于后续 PPO 更新。
4. 性能统计
generate_time 记录每次生成耗时,便于监控推理阶段的开销。
三、核心 PPO 更新流程(train_rlhf)
这是整个训练器的核心,实现了 策略网络(Actor)和价值网络(Critic)的同步更新,并包含关键的 GAE 优势计算、裁剪损失以及梯度溢出处理。
步骤 1:准备旧数据
从 inputs 中取出之前生成的 logprobs、ref_logprobs、rewards、values 等。同时根据 action_mask(即 attention_mask 去掉首 token)确定有效动作区间。
步骤 2:计算优势与回报(GAE)
get_advantages_and_returns 实现了 Generalized Advantage Estimation:
python
for t in reversed(range(start, length)):
delta = rewards[:, t] + gamma * nextvalues - values[:, t]
lastgaelam = delta + gamma * lam * lastgaelam
advantages_reversed.append(lastgaelam)
注意,我们会将超出对话结束位置(ends)的 reward 和 value 清零,确保后续计算不受填充 token 影响。优势函数 advantages 和回报 returns 均 detach,作为固定目标。
步骤 3:计算 Actor 损失并反传
当前策略重新前向 Actor 得到新的 actor_log_prob,然后调用 actor_loss_fn:
- 计算概率比
ratio = exp(logprobs - old_logprobs) - 应用 PPO 裁剪:
pg_loss = max(-advantages * ratio, -advantages * clip(ratio, 1-ε, 1+ε)) - 按 mask 求平均损失
调用 self.actor_model.backward(actor_loss) 完成梯度计算,但先不更新参数 (除非 align_overflow 为 False,则立即 step)。
步骤 4:计算 Critic 损失并反传
当前 Critic 前向得到新 values,调用 critic_loss_fn:
- 裁剪价值预测:
values_clipped = clip(values, old_values - ε, old_values + ε) - 损失为
0.5 * max((values - returns)^2, (values_clipped - returns)^2)的均值
同样调用 critic_model.backward(critic_loss)。
步骤 5:梯度溢出同步处理(align_overflow)
当启用 align_overflow 时,Actor 和 Critic 的优化器会分别检查梯度溢出(check_overflow)。由于二者独立进行梯度计算,可能出现一方溢出而另一方正常的情况。此时代码会强制跳过双方更新 (设置 skip_step = True),避免参数不同步导致训练崩溃。这是一种保守但稳健的策略,特别适用于混合精度训练。
最后,分别调用 actor_model.step() 和 critic_model.step() 完成参数更新(若溢出则跳过)。
步骤 6:返回损失值
返回 actor_loss 和 critic_loss 用于日志记录。
四、辅助功能与工程细节
1. 模型模式切换
train() / eval() 方法同时切换四个子模型的模式,确保在生成经验时所有模型均为评估模式(关闭 dropout 等),而在更新时 Actor 和 Critic 切换为训练模式,Ref 和 Reward 保持评估。
2. 模型范数监控
dump_model_norms 通过 get_model_norm 计算所有参数的全局范数,尤其适配 ZeRO-3 下的参数分片:对于状态为 NOT_AVAILABLE 的参数,会临时使用 GatheredParameters 上下文聚合,确保计算正确。这为调试和监控模型漂移提供了便利。
3. 无监督辅助损失
子类 DeepSpeedPPOTrainerUnsupervised 提供了 train_unsupervised 方法,允许在 RLHF 过程中混合无监督语言建模损失(如原始预训练任务),以缓解灾难性遗忘。该损失乘以系数 unsup_coef 后直接作用于 Actor 模型。
4. 浮点精度处理
当 compute_fp32_loss 为 True 时,会将 logits 转换为 float32 计算交叉熵,减小数值误差,尤其适用于 bf16 训练。
五、ZeRO-3 与分布式协调
代码中多处显式处理了 ZeRO-3 的特殊性:
- 生成时
synced_gpus=True确保所有 GPU 同步执行生成,避免因参数分片导致的死锁。 get_model_norm使用GatheredParameters临时收集分片参数。- 打印信息时使用
print_rank_0仅主节点输出,避免日志冗余。 - 通过
torch.distributed.all_reduce汇总各卡范数,实现全局监控。
此外,_validate_training_mode 等断言用于确保模式一致性,减少人为错误。
六、设计亮点总结
- 模块解耦:每个子模型独立管理优化器,训练器只负责调度 forward/backward/step,易于扩展和替换。
- 鲁棒性优先:无效回答缓存、溢出同步跳过、奖励裁剪等多重保险,使训练在数万步中保持稳定。
- 性能优化:一次前向获取 logits、value、reward 等所有信号,减少重复计算;利用 DeepSpeed 的梯度累积和 ZeRO 分片大幅降低显存占用。
- 可观测性:内置范数打印、生成时间统计、loss 输出,方便调试和调参。
七、使用注意事项
- 必须预先配置好每个模型的 DeepSpeed 引擎(包括 ZeRO stage、优化器类型、学习率等)。
- 输入 prompts 需包含结束符(如
end_of_conversation_token),且 tokenizer 对应正确。 - 对于 LLaMA 类模型,建议设置
do_sample=False或确保采样参数合理,避免 NaN。 - 若显存紧张,可减小
max_answer_seq_len或增大梯度累积步数。
八、完整代码
python
# Copyright (c) Microsoft Corporation.
# SPDX-License-Identifier: Apache-2.0
# DeepSpeed Team
import torch
import torch.nn.functional as F
import time
import deepspeed
from deepspeed.runtime.zero.partition_parameters import ZeroParamStatus
from deepspeed.accelerator import get_accelerator
from dschat.utils.utils import print_rank_0
def print_all_ranks(tag, value, rank):
world_size = torch.distributed.get_world_size()
all_tensor = torch.zeros(world_size, dtype=torch.float32).to(
get_accelerator().current_device_name())
all_tensor[rank] = value
torch.distributed.all_reduce(all_tensor, op=torch.distributed.ReduceOp.SUM)
print_rank_0(f'{tag} {all_tensor}', rank)
def get_model_norm(model):
with torch.no_grad():
total = 0.0
for param in model.parameters():
should_gather = hasattr(
param,
'ds_id') and param.ds_status == ZeroParamStatus.NOT_AVAILABLE
with deepspeed.zero.GatheredParameters(param,
enabled=should_gather):
total += float(param.float().norm())
return total
def gather_log_probs(logits, labels):
log_probs = F.log_softmax(logits, dim=-1)
log_probs_labels = log_probs.gather(dim=-1, index=labels.unsqueeze(-1))
return log_probs_labels.squeeze(-1)
class DeepSpeedPPOTrainer():
def __init__(self, rlhf_engine, args):
self.rlhf_engine = rlhf_engine
self.actor_model = self.rlhf_engine.actor
self.critic_model = self.rlhf_engine.critic
self.ref_model = self.rlhf_engine.ref
self.reward_model = self.rlhf_engine.reward
self.tokenizer = self.rlhf_engine.tokenizer
self.args = args
self.max_answer_seq_len = args.max_answer_seq_len
self.end_of_conversation_token_id = self.tokenizer(
args.end_of_conversation_token)['input_ids'][-1]
self.z3_enabled = args.actor_zero_stage == 3
self.compute_fp32_loss = self.args.compute_fp32_loss
# In case the generated experience is not valid (too short), we use the last valid
# generated experience. Alternatively, we can skip the step (on all workers).
# For now, use the last valid experience which is a simpler solution
self.last_generated_experience = None
# Those value can be changed
self.kl_ctl = 0.1
self.clip_reward_value = 5
self.cliprange = 0.2
self.cliprange_value = 0.2
self.gamma = 1.0
self.lam = 0.95
self.generate_time = 0.0
def _generate_sequence(self, prompts, mask, step):
max_min_length = self.max_answer_seq_len + prompts.shape[1]
# This has been added due to a probability/nan error that happens after
# meta-llama/Llama-2-7b-hf enabled do_sample:
# https://huggingface.co/meta-llama/Llama-2-7b-hf/commit/6fdf2e60f86ff2481f2241aaee459f85b5b0bbb9
if self.actor_model.module.config.model_type == "llama":
kwargs = dict(do_sample=False)
else:
kwargs = dict()
with torch.no_grad():
seq = self.actor_model.module.generate(
prompts,
attention_mask=mask,
max_length=max_min_length,
pad_token_id=self.tokenizer.pad_token_id,
synced_gpus=self.z3_enabled,
**kwargs)
# Filter out seq with no answers (or very short). This happens when users directly use the pre-training ckpt without supervised finetuning
# NOTE: this will causes each GPU has different number of examples
batch_size = seq.shape[0]
prompt_length = prompts.shape[1]
self.prompt_length = prompt_length
ans = seq[:, prompt_length:]
valid_ans_len = (ans != self.tokenizer.pad_token_id).sum(dim=-1)
if self.args.print_answers and (step % self.args.print_answers_interval
== 0):
print(
f"--- prompt --> step={step}, rank={torch.distributed.get_rank()}, {self.tokenizer.batch_decode(prompts, skip_special_tokens=True)}"
)
print(
f"--- ans --> step={step}, rank={torch.distributed.get_rank()}, {self.tokenizer.batch_decode(ans, skip_special_tokens=True)}"
)
out_seq = []
for i in range(batch_size):
if valid_ans_len[
i] <= 1: # if the answer is shorter than 1 token, drop it
print(
f'Dropping too short generated answer: {step=}: \n'
f'prompts: {self.tokenizer.batch_decode(prompts, skip_special_tokens=False)}\n'
f'answers: {self.tokenizer.batch_decode(ans, skip_special_tokens=False)}'
)
continue
else:
out_seq.append(seq[i:i + 1])
if not out_seq:
print(
f'All generated results are too short for rank={self.args.local_rank} step={step}\n'
f'-> prompts: {self.tokenizer.batch_decode(prompts, skip_special_tokens=False)}\n'
f'-> answers: {self.tokenizer.batch_decode(ans, skip_special_tokens=False)}'
)
return None
out_seq = torch.cat(out_seq, dim=0) # concat output in the batch dim
return out_seq
def generate_experience(self, prompts, mask, step):
self.eval()
generate_start = time.time()
seq = self._generate_sequence(prompts, mask, step)
generate_end = time.time()
if seq is None:
assert self.last_generated_experience is not None, f'Invalid generated experience at {step=}'
prompts = self.last_generated_experience['prompts']
seq = self.last_generated_experience['seq']
else:
self.last_generated_experience = {'prompts': prompts, 'seq': seq}
self.train()
pad_token_id = self.tokenizer.pad_token_id
attention_mask = seq.not_equal(pad_token_id).long()
with torch.no_grad():
output = self.actor_model(seq, attention_mask=attention_mask)
output_ref = self.ref_model(seq, attention_mask=attention_mask)
reward_score = self.reward_model.forward_value(
seq, attention_mask,
prompt_length=self.prompt_length)['chosen_end_scores'].detach(
)
values = self.critic_model.forward_value(
seq, attention_mask, return_value_only=True).detach()[:, :-1]
logits = output.logits
logits_ref = output_ref.logits
if self.compute_fp32_loss:
logits = logits.to(torch.float)
logits_ref = logits_ref.to(torch.float)
self.generate_time = generate_end - generate_start
return {
'prompts': prompts,
'logprobs': gather_log_probs(logits[:, :-1, :], seq[:, 1:]),
'ref_logprobs': gather_log_probs(logits_ref[:, :-1, :], seq[:,
1:]),
'value': values,
'rewards': reward_score,
'input_ids': seq,
"attention_mask": attention_mask
}
def compute_rewards(self, prompts, log_probs, ref_log_probs, reward_score,
action_mask):
kl_divergence_estimate = -self.kl_ctl * (log_probs - ref_log_probs)
rewards = kl_divergence_estimate
start = prompts.shape[1] - 1
ends = start + action_mask[:, start:].sum(1) + 1
reward_clip = torch.clamp(reward_score, -self.clip_reward_value,
self.clip_reward_value)
batch_size = log_probs.shape[0]
for j in range(batch_size):
rewards[j, start:ends[j]][-1] += reward_clip[j]
return rewards
def train_rlhf(self, inputs):
# train the rlhf mode here
### process the old outputs
prompts = inputs['prompts']
log_probs = inputs['logprobs']
ref_log_probs = inputs['ref_logprobs']
reward_score = inputs['rewards']
values = inputs['value']
attention_mask = inputs['attention_mask']
seq = inputs['input_ids']
start = prompts.size()[-1] - 1
action_mask = attention_mask[:, 1:]
old_values = values
with torch.no_grad():
old_rewards = self.compute_rewards(prompts, log_probs,
ref_log_probs, reward_score,
action_mask)
ends = start + action_mask[:, start:].sum(1) + 1
# we need to zero out the reward and value after the end of the conversation
# otherwise the advantage/return will be wrong
for i in range(old_rewards.shape[0]):
old_rewards[i, ends[i]:] = 0
old_values[i, ends[i]:] = 0
advantages, returns = self.get_advantages_and_returns(
old_values, old_rewards, start)
### process the new outputs
batch = {'input_ids': seq, "attention_mask": attention_mask}
actor_prob = self.actor_model(**batch, use_cache=False).logits
actor_log_prob = gather_log_probs(actor_prob[:, :-1, :], seq[:, 1:])
actor_loss = self.actor_loss_fn(actor_log_prob[:, start:],
log_probs[:, start:], advantages,
action_mask[:, start:])
self.actor_model.backward(actor_loss)
if not self.args.align_overflow:
self.actor_model.step()
value = self.critic_model.forward_value(**batch,
return_value_only=True,
use_cache=False)[:, :-1]
critic_loss = self.critic_loss_fn(value[:, start:], old_values[:,
start:],
returns, action_mask[:, start:])
self.critic_model.backward(critic_loss)
if self.args.align_overflow:
actor_overflow = self.actor_model.optimizer.check_overflow(
external=True)
critic_overflow = self.critic_model.optimizer.check_overflow(
external=True)
rank = torch.distributed.get_rank()
if actor_overflow and not critic_overflow:
self.critic_model.optimizer.skip_step = True
print_rank_0(
"OVERFLOW: actor overflow, skipping both actor and critic steps",
rank)
elif not actor_overflow and critic_overflow:
self.actor_model.optimizer.skip_step = True
print_rank_0(
"OVERFLOW: critic overflow, skipping both actor and critic steps",
rank)
elif actor_overflow and critic_overflow:
print_rank_0(
"OVERFLOW: actor and critic overflow, skipping both actor and critic steps",
rank)
self.actor_model.step()
self.critic_model.step()
return actor_loss, critic_loss
def get_overflow(self):
# Overflow is not expected when using bf16
# Therefore, DeepSpeed's BF16_Optimizer does not maintain an overflow indication
if self.args.dtype == "bf16":
return False, False
actor_overflow = self.actor_model.optimizer.overflow
critic_overflow = self.critic_model.optimizer.overflow
return actor_overflow, critic_overflow
def actor_loss_fn(self, logprobs, old_logprobs, advantages, mask):
## policy gradient loss
log_ratio = (logprobs - old_logprobs) * mask
ratio = torch.exp(log_ratio)
pg_loss1 = -advantages * ratio
pg_loss2 = -advantages * torch.clamp(ratio, 1.0 - self.cliprange,
1.0 + self.cliprange)
pg_loss = torch.sum(torch.max(pg_loss1, pg_loss2) * mask) / mask.sum()
return pg_loss
def critic_loss_fn(self, values, old_values, returns, mask):
## value loss
values_clipped = torch.clamp(
values,
old_values - self.cliprange_value,
old_values + self.cliprange_value,
)
if self.compute_fp32_loss:
values = values.float()
values_clipped = values_clipped.float()
vf_loss1 = (values - returns)**2
vf_loss2 = (values_clipped - returns)**2
vf_loss = 0.5 * torch.sum(
torch.max(vf_loss1, vf_loss2) * mask) / mask.sum()
return vf_loss
def get_advantages_and_returns(self, values, rewards, start):
# Adopted from https://github.com/CarperAI/trlx/blob/main/trlx/models/modeling_ppo.py#L134
lastgaelam = 0
advantages_reversed = []
length = rewards.size()[-1]
for t in reversed(range(start, length)):
nextvalues = values[:, t + 1] if t < length - 1 else 0.0
delta = rewards[:, t] + self.gamma * nextvalues - values[:, t]
lastgaelam = delta + self.gamma * self.lam * lastgaelam
advantages_reversed.append(lastgaelam)
advantages = torch.stack(advantages_reversed[::-1], dim=1)
returns = advantages + values[:, start:]
return advantages.detach(), returns
def _validate_training_mode(self):
assert self.actor_model.module.training
assert self.critic_model.module.training
def _validate_evaluation_mode(self):
assert not self.actor_model.module.training
assert not self.critic_model.module.training
assert not self.ref_model.module.training
assert not self.reward_model.module.training
def train(self):
self.actor_model.train()
self.critic_model.train()
def eval(self):
self.actor_model.eval()
self.critic_model.eval()
self.reward_model.eval()
self.ref_model.eval()
def dump_model_norms(self, tag):
actor_model_norm = get_model_norm(self.actor_model)
ref_model_norm = get_model_norm(self.ref_model)
critic_model_norm = get_model_norm(self.critic_model)
reward_model_norm = get_model_norm(self.reward_model)
print_all_ranks(f'{tag} global_actor_model_norm', actor_model_norm,
self.args.local_rank)
print_all_ranks(f'{tag} global_ref_model_norm', ref_model_norm,
self.args.local_rank)
print_all_ranks(f'{tag} global_critic_model_norm', critic_model_norm,
self.args.local_rank)
print_all_ranks(f'{tag} global_reward_model_norm', reward_model_norm,
self.args.local_rank)
class DeepSpeedPPOTrainerUnsupervised(DeepSpeedPPOTrainer):
def __init__(self, *args, **kwargs):
super().__init__(*args, **kwargs)
def train_unsupervised(self, inputs, unsup_coef):
# Train the unsupervised model here
self._validate_training_mode()
outputs = self.actor_model(**inputs, use_cache=False)
loss = outputs.loss
self.actor_model.backward(unsup_coef * loss)
self.actor_model.step()
return loss