Prhub

#1882 fix ppo value offload bugs

原始 PR 作者 lilei199908 合并时间 2026-05-06 12:39 文件变更 8 提交数 30 评论 2 代码增减 +181 / -9

执行摘要

修复 PPO 使用 critic 时的 offload 训练 bug

PR标题和description为空,但从代码变更推断,原有PPO训练逻辑在use_critic=True搭配offload_train时存在 actor 进程休眠后未正确断开rollout引擎NCCL通信组、唤醒后未重建连接的问题,导致权重更新失败或程序崩溃。本修复旨在确保critic存在时actor的offload行为符合预期,避免通信状态不一致。

建议精读。本PR修复了PPO多角色offload训练中的关键状态管理问题,展示了异步训练中actor与rollout引擎连接维护的典型解决方案:sleep前断开、wake后重建。design_pattern值得参考,特别是disconnect/connect的显式化。另外参数强制绑定的策略也值得讨论——究竟是缺陷修复还是约束收紧,不同团队可能有不同偏好。

讨论亮点

Review中Copilot bot指出两个assert安全问题:sleep()wake_up()内部assert self.args.offload_train,但调用条件已成为offload_train or use_critic,若use_critic=Trueoffload_train=False(尽管参数校验已强制绑定,但逻辑上仍可能被绕过),assert会失败。作者通过参数校验的强制绑定解决了该隐患,但assert仍保留,依赖上层校验保障安全性。

实现拆解

  1. 强制offload_train:在slime/utils/arguments.pyslime_validate_args中,当use_critic=True时自动设置offload_train=True,保证后续逻辑前提条件。
  2. 添加disconnect_rollout_engines方法:在slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py中新增公开方法,用于显式断开rollout引擎的NCCL连接并清空model_update_groups
  3. 调整actor休眠与唤醒:在slime/backends/megatron_utils/actor.pysleep方法中,增加条件:当角色为actor、使用critic、非colocate且weight_updater有disconnect_rollout_engines时,在销毁进程组前主动断开连接;在train入口增加条件:当offload_train或use_critic时先唤醒;在train结尾增加相应条件睡眠;在update_weights中也增加重连逻辑:如果需重连则先wake_up,连接rollout引擎,完成更新后睡眠。
  4. 简化GPU偏移计算:在slime/backends/sglang_utils/sglang_engine.pyget_base_gpu_id中移除已废弃的use_critic分支,因为actor和critic不再混排GPU分配。
  5. 调整训练入口:在train.py中确保update_weightsoffload_rollout解耦,使其在offload_rollout之外也能正确执行。
  6. 新增集成测试tests/test_qwen3_4B_ppo_disaggregate.py涵盖Qwen3-4B在disaggregate模式下使用critic的端到端PPO训练流程,通过条件变量ENABLE_EVALTIGHT_HOST_MEMORY控制评估和内存压力。
文件 模块 状态 重要度
slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py 权重同步 modified 6.52
slime/backends/megatron_utils/actor.py Actor 模块 modified 6.44
tests/test_qwen3_4B_ppo_disaggregate.py PPO 解耦 added 6.85
slime/backends/sglang_utils/sglang_engine.py SGLang 引擎 modified 5.37
slime/utils/arguments.py 参数验证 modified 5.16
train.py 训练入口 modified 4.32
.github/workflows/pr-test.yml CI 配置 modified 2.55
.github/workflows/pr-test.yml.j2 CI 配置 modified 2.24

关键符号

disconnect_rollout_engines sleep wake_up update_weights get_base_gpu_id slime_validate_args prepare execute

关键源码片段

slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py core-logic

新增 disconnect_rollout_engines 方法,提供显式断开 rollout 引擎 NCCL 通信组的功能,是修复连接管理的关键。

class WeightUpdater:
    # ... 其他方法 ...
​
    def disconnect_rollout_engines(self) -> None:
        """如果当前是PP源角色且存在model_update_groups,则安全断开与rollout引擎的NCCL连接"""
        if not getattr(self, "_is_pp_src_rank", False) or self._model_update_groups is None:
            return
        # 调用底层函数断开连接
        disconnect_rollout_engines_from_distributed(
            self.args, self._group_name, self._model_update_groups, self.rollout_engines
        )
        self._model_update_groups = None # 置空标志,防止重复断开
​
    @torch.no_grad()
    def update_weights(self) -> None:
        """权重更新主循环:暂停rollout引擎、清缓存、广播非专家参数、广播专家参数、恢复引擎"""
        self.weight_version += 1
        # ... 具体广播逻辑 ...
slime/backends/megatron_utils/actor.py core-logic

核心训练文件,修改 sleep/wake_up/update_weights 等方法的条件控制,确保 actor 在 critic 模式下正确管理 rollout 引擎连接。

@timer
def sleep(self) -> None:
    """挂起actor,释放GPU资源;若使用critic则先断开rollout引擎连接"""
    assert self.args.offload_train
    clear_memory(clear_host_memory=True)
    print_memory("before offload model")
    # 需要断开连接的条件:是 actor、使用 critic、非 colocate、weight_updater 支持断开
    if (
        self.role == "actor"
        and self.args.use_critic
        and not self.args.colocate
        and hasattr(self.weight_updater, "disconnect_rollout_engines")
    ):
        self.weight_updater.disconnect_rollout_engines()
    destroy_process_groups()
    torch_memory_saver.pause()
    print_memory("after offload model")
​
​
# 在 update_weights 中的关键分支(在更新权重前后)
# 判断是否需要完整的重连流程
reconnect_rollout_engines = self.args.offload_train and self.args.use_critic and not self.args.colocateif reconnect_rollout_engines:
    self.wake_up() # 唤醒 actor 进程组
elif self.args.offload_train:
    reload_process_groups() # 仅重载进程组# 如果有新引擎加入或需要重连,则连接 rollout 引擎
if num_new_engines > 0 or reconnect_rollout_engines:
    self.weight_updater.connect_rollout_engines(
        rollout_engines, rollout_engine_lock,
    )# ... 权重更新步骤 ...# 更新完成后,如果之前唤醒过则再次休眠
if reconnect_rollout_engines:
    self.sleep()
elif self.args.offload_train:
    destroy_process_groups()

评论区精华

sleep/wake_up 断言与调用条件不一致 正确性

Copilot 指出 sleep() 和 wake_up() 内部 assert self.args.offload_train,但调用条件变为 offload_train or use_critic,若 use_critic=True 且 offload_train=False 将崩溃。

结论:作者通过 slime_validate_args 强制在 use_critic 时设置 offload_train=True,避免了断言触发的可能性,但 assert 仍存在,依赖参数校验的约束。 · 已解决

风险与影响

  • 强制修改 use_critic 时 offload_train=True,可能改变用户习惯,但实际训练中模型卸载是常见需求,风险较低。
  • sleep/wake_up 中的 assert 仍保留,若有人绕过参数校验直接调用底层接口,仍可能触发。
  • 新添加的 disconnect_rollout_engines 在 colocate 模式下不会被调用,但 colocate 与 critic 共存场景未明确测试。
  • 测试仅覆盖了 disaggregate 场景,colocate 或纯 actor 场景未新增测试。
  • 用户:开启critic训练的PPO用户在offload_train场景下不再遇到连接丢失崩溃;强制开启offload_train可能增加显存压力,但原本critic训练通常也需要offload。
  • 系统:修改了核心训练循环(actor.py train),影响所有使用critic的PPO训练。
  • 团队:需要关注后续版本中critic与offload的耦合关系是否需解耦。
  • 影响程度:中,涉及训练流程正确性。
critic 与 offload 强制绑定 sleep/wake_up 断言隐患 colocate 模式未测试 仅 disaggregate 测试覆盖

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论