深入解析 DeepSpeedPPOTrainer:基于 PPO 的大规模 RLHF 训练实现

在大语言模型(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 四个模型并行前向,分别得到:

  • logprobsref_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 中取出之前生成的 logprobsref_logprobsrewardsvalues 等。同时根据 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 等断言用于确保模式一致性,减少人为错误。


六、设计亮点总结

  1. 模块解耦:每个子模型独立管理优化器,训练器只负责调度 forward/backward/step,易于扩展和替换。
  2. 鲁棒性优先:无效回答缓存、溢出同步跳过、奖励裁剪等多重保险,使训练在数万步中保持稳定。
  3. 性能优化:一次前向获取 logits、value、reward 等所有信号,减少重复计算;利用 DeepSpeed 的梯度累积和 ZeRO 分片大幅降低显存占用。
  4. 可观测性:内置范数打印、生成时间统计、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
相关推荐
05664644 分钟前
python高级——Python 类型提示与 Pydantic
网络·人工智能·windows·python·学习
小女孩真可爱44 分钟前
GPT(3)----------------GQA分组查询注意力机制提速
人工智能·pytorch·gpt·深度学习·大模型
qq_25294131681 小时前
山体滑坡目标检测数据集 | 山体滑坡检测 地质灾害识别 遥感监测 目标检测 YOLO格式
人工智能·yolo·目标检测·计算机视觉·视觉检测·自然灾害·滑坡数据集
tachibana21 小时前
大语言模型基础
数据库·人工智能·语言模型·自然语言处理·大模型·llm
旋转的油纸伞1 小时前
Wukong: Towards a Scaling Law for Large-Scale Recommendation
人工智能·深度学习·神经网络·目标检测·机器学习·自然语言处理·caffe
硅基流动1 小时前
山东铁路基金公司与硅基流动达成战略合作,共建 Token 工厂
人工智能·科技
飞哥数智坊1 小时前
难道 AI 真要让程序员三班倒了?
人工智能·ai编程
玫瑰互动GEO1 小时前
抖音SEO优化技术拆解:搜索排名四大因子与4步落地算法分析
人工智能·算法·搜索引擎·语音识别
IT_陈寒1 小时前
Vite打包时踩了个坑,static资源去哪了?
前端·人工智能·后端