执行摘要
- 一句话:修复PPO使用critic时的offload训练bug
- 推荐动作:建议精读。本PR修复了PPO多角色offload训练中的关键状态管理问题,展示了异步训练中actor与rollout引擎连接维护的典型解决方案:sleep前断开、wake后重建。design_pattern值得参考,特别是disconnect/connect的显式化。另外参数强制绑定的策略也值得讨论——究竟是缺陷修复还是约束收紧,不同团队可能有不同偏好。
功能与动机
PR标题和description为空,但从代码变更推断,原有PPO训练逻辑在use_critic=True搭配offload_train时存在 actor 进程休眠后未正确断开rollout引擎NCCL通信组、唤醒后未重建连接的问题,导致权重更新失败或程序崩溃。本修复旨在确保critic存在时actor的offload行为符合预期,避免通信状态不一致。
实现拆解
- 强制offload_train:在
slime/utils/arguments.py的slime_validate_args中,当use_critic=True时自动设置offload_train=True,保证后续逻辑前提条件。
- 添加disconnect_rollout_engines方法:在
slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py中新增公开方法,用于显式断开rollout引擎的NCCL连接并清空model_update_groups。
- 调整actor休眠与唤醒:在
slime/backends/megatron_utils/actor.py的sleep方法中,增加条件:当角色为actor、使用critic、非colocate且weight_updater有disconnect_rollout_engines时,在销毁进程组前主动断开连接;在train入口增加条件:当offload_train或use_critic时先唤醒;在train结尾增加相应条件睡眠;在update_weights中也增加重连逻辑:如果需重连则先wake_up,连接rollout引擎,完成更新后睡眠。
- 简化GPU偏移计算:在
slime/backends/sglang_utils/sglang_engine.py的get_base_gpu_id中移除已废弃的use_critic分支,因为actor和critic不再混排GPU分配。
- 调整训练入口:在
train.py中确保update_weights与offload_rollout解耦,使其在offload_rollout之外也能正确执行。
- 新增集成测试:
tests/test_qwen3_4B_ppo_disaggregate.py涵盖Qwen3-4B在disaggregate模式下使用critic的端到端PPO训练流程,通过条件变量ENABLE_EVAL和TIGHT_HOST_MEMORY控制评估和内存压力。
关键文件:
slime/backends/megatron_utils/update_weight/update_weight_from_distributed.py(模块 权重同步;类别 source;类型 core-logic;符号 disconnect_rollout_engines): 新增disconnect_rollout_engines方法,提供显式断开rollout引擎NCCL通信组的功能,是修复连接管理的关键。
slime/backends/megatron_utils/actor.py(模块 Actor模块;类别 source;类型 core-logic;符号 init, sleep, wake_up, update_weights): 核心训练文件,修改sleep/wake_up/update_weights等方法的条件控制,确保actor在critic模式下正确管理rollout引擎连接。
tests/test_qwen3_4B_ppo_disaggregate.py(模块 PPO解耦;类别 test;类型 test-coverage;符号 prepare, execute): 新增PPO disaggregate模式端到端测试,验证修复后的critic+offload训练正确性。
slime/backends/sglang_utils/sglang_engine.py(模块 SGLang引擎;类别 source;类型 core-logic;符号 get_base_gpu_id): 移除已废弃的use_critic GPU偏移计算分支,简化逻辑。
slime/utils/arguments.py(模块 参数验证;类别 source;类型 core-logic;符号 slime_validate_args): 强制offload_train与use_critic绑定,保证后续逻辑的前提条件。
train.py(模块 训练入口;类别 source;类型 core-logic): 调整update_weights调用时机,确保在offload_rollout以外的场景也能执行。
.github/workflows/pr-test.yml(模块 CI配置;类别 infra;类型 infrastructure): CI配置的小调整,增加新测试触发。
.github/workflows/pr-test.yml.j2(模块 CI配置;类别 infra;类型 infrastructure): CI模板调整,配合新增测试。
关键符号: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
新增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
核心训练文件,修改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.colocate
if 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()
评论区精华
Review中Copilot bot指出两个assert安全问题:sleep()和wake_up()内部assert self.args.offload_train,但调用条件已成为offload_train or use_critic,若use_critic=True而offload_train=False(尽管参数校验已强制绑定,但逻辑上仍可能被绕过),assert会失败。作者通过参数校验的强制绑定解决了该隐患,但assert仍保留,依赖上层校验保障安全性。
- sleep/wake_up断言与调用条件不一致 (correctness): 作者通过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测试覆盖
关联脉络
- PR #1856 refactor/ppo: PPO训练架构重构,解耦actor-critic配置与通信,本PR可能依赖于该重构引入的配置和通信机制。
- PR #1878 fix ppo value head load bugs: 同一系列的PPO bug修复,涉及value head加载,与本PR的offload bug修复同属PPO训练稳定性改进。
参与讨论