强化学习训练最让人血压飙升的场景,不是reward不涨,而是跑了十几个小时的训练任务在半夜挂掉,第二天早上发现checkpoint还是六个小时前的。尤其是现在RL框架普遍采用分离式架构——训练侧和推理侧各自独立部署,参数同步、权重更新、状态保存这几件事一旦没有统一的协调机制,故障恢复就变成了一场灾难。这篇文章想聊的,就是怎么把Checkpoint Engine这套东西接进RL框架里,让训练任务从"能跑"进化到"挂了也能接着跑"。
我所在的团队最近在做一个中等规模的RLHF训练项目,推理侧用的是SGLang做rollout生成,训练侧是自研的PyTorch FSDP框架,中间靠Parameter Server做权重同步。项目初期一切顺利,直到有一次集群网络抖动导致推理节点集体掉线,训练侧还在傻等rollout结果,等我们发现的时候已经空转了四十分钟。更糟的是,由于checkpoint保存逻辑和参数同步逻辑是两套独立的东西,恢复的时候权重版本对不上,只能从头再来。那次事故之后,我们下定决心把Checkpoint Engine作为一等公民接入整个RL训练链路。
这篇文章适合正在搭建或维护RL训练框架的工程师,尤其是那些已经踩过或即将踩到"故障恢复"这个坑的人。我会从架构设计的角度讲清楚Checkpoint Engine在RL场景下和传统训练有什么不同,然后给出具体的接入方案、参数同步的协调逻辑、故障恢复的完整流程,最后分享几个我们在实操中踩过的坑和对应的解法。读者不需要对SGLang或Parameter Server有深入了解,但最好对RL训练的基本流程有概念。
1. 为什么RL框架的Checkpoint比传统训练复杂得多
1.1 传统训练Checkpoint的隐含假设在RL场景下全部失效
传统监督学习训练的checkpoint逻辑相对简单:模型参数、优化器状态、学习率调度器的状态,再加上当前的step数,打包存下来就完事了。恢复的时候把这些东西load回去,训练就能无缝继续。这套逻辑成立的前提是:训练过程中只有一份状态在变化,而且这份状态完全由训练进程自己掌控。
RL训练打破了这个前提。在典型的RLHF或RLVR流程里,至少存在三份需要协调的状态:训练侧的策略模型参数(包括优化器状态)、推理侧的生成模型权重(用于rollout)、以及rollout产生的经验数据(experience buffer)。这三份状态分别由不同的进程甚至不同的节点管理,它们之间通过参数同步机制保持一致性。当你需要保存checkpoint的时候,问题就来了:你保存的是哪一份状态?三份状态之间的版本关系怎么记录?恢复的时候怎么保证三份状态回到同一个一致的时间点?
我们最初的做法是只保存训练侧的checkpoint,推理侧的权重在恢复后重新从训练侧同步一次。听起来很合理对吧?但实际跑起来发现,如果checkpoint保存的时候参数同步正在进行中,训练侧保存的权重版本和推理侧实际使用的权重版本可能差了好几个step。恢复之后,推理侧用新权重生成rollout,但experience buffer里可能还残留着旧权重生成的数据,导致训练信号混乱。
1.2 参数同步的中间态是故障恢复的最大敌人
Parameter Server架构下,权重同步通常是一个异步过程。训练侧算完梯度、更新完参数之后,把新权重推送到Parameter Server,推理侧再从Parameter Server拉取最新权重。这个"推送-拉取"的窗口期内,系统处于一个中间态:训练侧已经是新版本了,推理侧可能还是旧版本,Parameter Server上的版本可能介于两者之间。
如果在这个窗口期内发生故障,checkpoint的状态就是不确定的。我们遇到过好几次这样的情况:故障恢复后训练loss突然飙升,排查了半天才发现是推理侧加载了一个"半新不旧"的权重版本,生成的rollout质量急剧下降,训练侧基于这些差数据更新参数,直接把模型带偏了。
解决这个问题的核心思路是:checkpoint必须记录一个全局一致的版本号,这个版本号要能唯一标识"训练侧参数、推理侧权重、experience buffer"三者的一致状态。任何时刻,只有当一个完整的同步周期结束之后,才能触发checkpoint保存。换句话说,checkpoint的保存点必须落在同步周期的边界上,而不能落在中间。
1.3 推理侧状态的特殊性:KV Cache和生成配置
还有一个容易被忽略的点:推理侧不只是模型权重需要保存。SGLang这类推理引擎在运行时会维护大量的KV Cache和请求级别的生成状态。虽然这些状态在故障恢复后可以重建,但重建的成本很高,尤其是当rollout任务队列很长的时候。
我们的做法是把推理侧的状态分成两类:必须持久化的和可以重建的。必须持久化的是模型权重版本号和生成配置(比如temperature、top_p这些采样参数),这些信息决定了rollout的语义一致性。可以重建的是KV Cache和请求队列,故障恢复后重新初始化即可,但需要在恢复流程中显式地清空这些状态,避免残留数据污染新的rollout。
2. Checkpoint Engine的接入点选择与架构设计
2.1 三种接入方案的对比:训练侧主导、推理侧主导、独立协调器
把Checkpoint Engine接入RL框架,最核心的架构决策是:谁来触发checkpoint保存?谁来管理版本号?谁来协调恢复流程?我们评估了三种方案,各有优劣。
第一种是训练侧主导。训练进程在每N个step之后触发checkpoint保存,同时通知推理侧暂停rollout,等推理侧确认当前同步周期结束后,训练侧保存自己的状态,并记录当前的全局版本号。这种方案实现简单,但问题是训练侧需要等待推理侧确认,会阻塞训练流程。如果推理侧响应慢,训练效率会明显下降。
第二种是推理侧主导。推理侧在完成一批rollout之后触发checkpoint,通知训练侧保存状态。这种方案适合rollout时间远长于训练step时间的场景,但同样存在协调开销。
第三种是独立协调器。单独起一个轻量级的协调进程,它不参与实际的训练和推理计算,只负责管理全局版本号和触发checkpoint。训练侧和推理侧都向协调器注册自己的状态,协调器在确认所有参与方都到达同步边界后,统一触发保存。这种方案解耦最彻底,但引入了一个额外的组件,增加了部署复杂度。
我们最终选择了第三种方案,原因是我们的训练和推理是跨节点部署的,训练侧和推理侧的网络延迟不稳定,如果让一方等待另一方,很容易出现长尾延迟。独立协调器可以用心跳机制异步地收集各方的状态,只在真正需要保存checkpoint的时候才做同步等待。
2.2 全局版本号的设计:单调递增还是向量时钟
全局版本号的设计直接决定了恢复逻辑的复杂度。最简单的方案是单调递增的整数版本号,每次完整的同步周期结束后加一。这种方案的好处是直观、易于比较,但缺点是丢失了版本之间的因果关系。比如版本号从5跳到7,你只知道中间发生了一次同步,但不知道这次同步涉及哪些参数的变更。
我们一开始用的就是单调递增版本号,后来发现一个问题:当训练侧和推理侧的同步频率不一致时(训练侧每个step都更新参数,推理侧每K个step才拉取一次权重),单调版本号无法表达"推理侧当前使用的权重对应训练侧的哪个step"。这导致恢复的时候经常出现版本对不上的情况。
后来我们改成了向量时钟的方案:每个参与方维护自己的本地版本号,全局版本号是一个向量(training_step, inference_version, buffer_version)。checkpoint保存的时候记录这个向量,恢复的时候要求所有参与方都回到向量指定的状态。这种方案的缺点是版本比较逻辑更复杂,但好处是恢复的精确度大大提高。实际实现中,我们用了一个简化的向量时钟:只记录训练侧的step数和推理侧的权重版本号,experience buffer的版本号通过关联训练step来推导。
2.3 协调器的心跳机制与超时处理
独立协调器的核心是一个心跳收集循环。训练侧和推理侧每隔固定间隔(我们用的是5秒)向协调器发送心跳,心跳里包含当前的本地版本号和状态标识(idle、syncing、rollout_in_progress等)。协调器维护一个状态表,记录每个参与方的最后心跳时间和当前状态。
当协调器决定触发checkpoint时(可以基于时间间隔、step数或者手动触发),它会向所有参与方发送"准备checkpoint"的信号。参与方收到信号后,需要完成当前正在进行的原子操作(比如训练侧完成当前step的参数更新,推理侧完成当前batch的rollout),然后进入"ready"状态并通知协调器。协调器等待所有参与方都进入"ready"状态后,再发送"执行checkpoint"的信号。参与方执行保存操作,完成后通知协调器。协调器收到所有确认后,更新全局版本号,并通知参与方恢复正常的训练和推理流程。
超时处理是必须的。如果某个参与方在指定时间内没有响应"准备checkpoint"信号(我们设置的超时是30秒),协调器会放弃本次checkpoint,记录一条警告日志,并让已经进入"ready"状态的参与方恢复运行。这里有个细节:放弃checkpoint之后,全局版本号不能变,否则会导致版本号空洞。我们最初没注意这一点,结果出现了一次版本号跳跃,恢复的时候找不到对应的checkpoint文件。
3. 参数同步与Checkpoint保存的协调逻辑
3.1 同步周期的边界定义:什么时刻才算"一致状态"
参数同步的边界定义是整个协调逻辑的基础。在我们的架构里,一个完整的同步周期包含以下步骤:训练侧完成一个step的参数更新,将新权重推送到Parameter Server,Parameter Server确认接收并更新版本号,推理侧从Parameter Server拉取最新权重,推理侧确认加载完成并开始使用新权重生成rollout。
只有当推理侧确认开始使用新权重之后,系统才进入一致状态。在此之前,任何时刻保存的checkpoint都可能包含不一致的状态。所以我们的协调器只在收到推理侧的"新权重已生效"确认之后,才认为当前同步周期结束,才允许触发checkpoint。
这里有个性能优化的点:如果每个step都等待推理侧确认,训练效率会很低。我们的做法是允许训练侧连续进行多个step的参数更新,但推理侧的权重拉取是异步的。协调器记录训练侧的最新step数和推理侧的最新权重版本号,当两者的差距超过阈值(我们设置的阈值是5个step)时,协调器会通知训练侧暂停,等待推理侧追上。这样既保证了训练效率,又避免了版本差距过大导致的恢复困难。
3.2 Checkpoint保存时的参数冻结策略
保存checkpoint的时候,必须冻结参数更新,否则保存出来的状态可能是不一致的。但冻结的时间不能太长,否则会影响训练效率。我们的策略是"最小冻结窗口":协调器发送"准备checkpoint"信号后,训练侧完成当前step的参数更新就立即暂停,不再开始新的step。推理侧完成当前batch的rollout后暂停,不再接受新的rollout请求。
从发送信号到所有参与方进入"ready"状态,我们的实测平均耗时是2-3秒,最坏情况下(推理侧正在处理一个长序列的rollout)可能达到10秒以上。为了减少这个时间,我们把推理侧的rollout batch大小控制在一个合理的范围内,避免单个batch的生成时间过长。同时,训练侧的step时间也要控制,如果单个step的耗时超过5秒,就需要考虑减小batch size或者优化计算图。
参数冻结期间,Parameter Server仍然可以接受读请求(推理侧可能需要读取权重来完成当前的rollout),但不能接受写请求(训练侧不能推送新权重)。我们在Parameter Server的接口层面加了一个"冻结标志",当协调器触发checkpoint时,Parameter Server设置这个标志,拒绝所有写请求,直到checkpoint完成。
3.3 保存内容的分层设计:哪些必须存,哪些可以重建
Checkpoint的保存内容需要分层设计,不能什么都存,也不能什么都不存。我们的分层方案是这样的:
第一层是必须持久化的核心状态,包括训练侧的策略模型参数、优化器状态、学习率调度器状态、当前的训练step数,以及推理侧的模型权重版本号和生成配置。这些状态决定了训练和推理的语义一致性,丢失任何一项都会导致恢复后的行为不一致。
第二层是建议持久化的辅助状态,包括experience buffer中的样本数据、rollout的统计信息(比如平均reward、生成长度分布)、训练侧的梯度累积状态。这些状态在恢复后可以通过重新生成或重新计算来重建,但重建成本较高,所以建议持久化。
第三层是可以重建的临时状态,包括推理侧的KV Cache、请求队列、训练侧的数据加载器状态。这些状态在恢复后重新初始化即可,不需要持久化。但需要注意的是,恢复流程中必须显式地清空这些状态,避免残留数据污染新的训练过程。
存储方面,我们用的是分层的存储策略:核心状态存在高可用的分布式文件系统上,保证随时可读;辅助状态存在本地SSD上,定期同步到远程;临时状态不存储。这样既保证了恢复的可靠性,又控制了存储成本。
4. 故障恢复的完整流程与版本一致性校验
4.1 故障检测:怎么判断是"真故障"还是"慢节点"
故障恢复的第一步是准确地检测故障。在分布式训练中,最棘手的问题不是检测到节点掉线,而是区分"真故障"和"慢节点"。我们遇到过好几次这样的情况:某个推理节点因为负载过高,心跳延迟了几十秒,协调器误判为故障,触发了不必要的恢复流程,结果恢复过程中那个"慢节点"又活过来了,导致状态混乱。
我们的解决方案是引入"疑似故障"和"确认故障"两级判定。当某个参与方的心跳超时超过阈值(我们设置的是15秒)时,协调器将其标记为"疑似故障",并开始一个观察窗口(我们设置的是60秒)。在观察窗口内,如果心跳恢复,则取消疑似标记;如果心跳持续缺失,则升级为"确认故障",触发恢复流程。
观察窗口的长度需要根据实际部署环境来调整。如果网络抖动比较频繁,观察窗口可以设长一些;如果对恢复速度要求高,可以设短一些但配合更灵敏的心跳机制。我们最终用的是15秒心跳间隔加60秒观察窗口的组合,在实际运行中误判率很低。
4.2 恢复时的版本回滚:如何找到最近的"一致checkpoint"
确认故障后,恢复流程的第一步是找到最近的"一致checkpoint"。这里的关键是:不是所有保存的checkpoint都是一致的。如果checkpoint保存过程中发生了故障,可能留下一个不完整的checkpoint文件。我们的做法是在checkpoint保存完成后,写入一个"完成标记"文件,恢复的时候只加载带有完成标记的checkpoint。
找到最近的完成checkpoint后,需要校验版本一致性。协调器读取checkpoint中记录的全局版本号向量,然后检查当前存活的参与方的本地版本号。如果某个参与方的本地版本号高于checkpoint中的版本号,说明它在checkpoint之后又进行了更新,需要回滚到checkpoint指定的版本。如果低于,说明它落后了,需要从checkpoint中恢复状态。
版本回滚的具体操作取决于参与方的类型。训练侧的回滚相对简单:加载checkpoint中的模型参数和优化器状态,覆盖当前状态即可。推理侧的回滚需要额外注意:不仅要加载checkpoint中的权重版本,还要清空KV Cache和请求队列,确保没有残留的旧状态。
4.3 恢复后的参数重新同步:避免"版本悬崖"
恢复完成后,训练侧和推理侧的版本号可能不一致(比如训练侧回滚到了step 100,推理侧回滚到了权重版本50,但训练侧的最新step是100,推理侧需要拉取权重版本100)。这时候需要触发一次参数重新同步,让推理侧追上训练侧的版本。
这个重新同步的过程需要特别注意,我们称之为"版本悬崖"问题:如果训练侧在恢复后立即开始新的step,而推理侧还在同步旧版本,那么推理侧生成的rollout可能基于过时的权重,导致训练信号不一致。我们的做法是:恢复完成后,训练侧先暂停,等待推理侧完成权重同步并确认新权重生效,然后再恢复训练。这个等待时间通常很短(几秒到几十秒),但能有效避免版本不一致的问题。
还有一个细节:恢复后的第一次参数同步,我们强制使用全量同步而不是增量同步。虽然全量同步的开销更大,但它能确保推理侧的权重和训练侧完全一致,避免增量同步可能带来的累积误差。在后续的正常训练中,再切换回增量同步。
5. 实操中踩过的坑与对应的解法
5.1 checkpoint文件写入过程中的节点掉线
这是我们在早期遇到的最频繁的问题。checkpoint文件通常比较大(我们的模型参数加上优化器状态大概有几十GB),写入过程需要几十秒甚至几分钟。如果在这个期间节点掉线,checkpoint文件就会处于不完整状态。更糟糕的是,如果协调器没有正确检测到写入失败,可能会把这个不完整的checkpoint标记为"完成",导致恢复时加载了损坏的数据。
我们的解法是采用"临时文件+原子重命名"的写入策略。checkpoint先写入一个临时文件(比如checkpoint_100.tmp),写入完成后计算文件的校验和,然后将临时文件重命名为正式文件(checkpoint_100.pt),同时写入一个包含校验和的元数据文件。恢复的时候,先检查元数据文件是否存在且校验和匹配,只有都满足才加载checkpoint。这个策略看起来简单,但实际效果非常好,自从采用之后再也没有出现过加载损坏checkpoint的情况。
5.2 推理侧权重版本与训练侧step数的映射错乱
这个问题比较隐蔽,我们排查了很久才找到根因。现象是:故障恢复后,训练loss正常下降了一段时间,然后突然飙升。检查日志发现,推理侧使用的权重版本和训练侧的step数之间的映射关系错乱了。比如训练侧在step 150,但推理侧加载的是对应step 120的权重。
根因在于我们的版本号映射逻辑:训练侧每推送一次权重到Parameter Server,Parameter Server的版本号加一。但训练侧并不是每个step都推送权重(我们设置的是每3个step推送一次),所以Parameter Server的版本号和训练侧的step数之间是一个非线性的映射关系。故障恢复的时候,我们直接用Parameter Server的版本号去推算训练侧的step数,结果算错了。
解法是维护一个显式的映射表,记录每次权重推送时的(训练step数, Parameter Server版本号)对。这个映射表也需要持久化到checkpoint中,恢复的时候直接查表,而不是通过计算来推算。这个映射表的数据量很小(每个条目几十字节),对存储和传输的开销可以忽略不计。
5.3 恢复后experience buffer的样本分布偏移
这个问题是在一次大规模故障恢复后发现的。恢复之后,训练loss虽然能下降,但策略的熵值明显偏低,生成的回复多样性下降。排查后发现,experience buffer中的样本分布出现了偏移:恢复前buffer中积累了大量来自旧策略的样本,恢复后新策略生成的样本还没有充分填充buffer,导致训练时旧样本的权重过高。
解法是在恢复流程中加入buffer的"新鲜度检查"。具体来说,我们给buffer中的每个样本打上一个"生成时的策略版本号"标签。恢复后,如果buffer中超过一定比例(我们设置的阈值是30%)的样本来自过时的策略版本,就清空buffer,重新生成样本。这个策略会浪费一些计算资源,但能保证训练信号的准确性。后来我们进一步优化,不是全部清空,而是对过时样本进行降权处理,在计算梯度时给它们一个较小的权重。
5.4 协调器单点故障的隐患
独立协调器方案虽然解耦彻底,但引入了一个单点故障。如果协调器本身挂掉,整个训练任务就会失去协调能力,无法触发checkpoint,也无法进行故障恢复。我们最初的实现没有考虑这个问题,直到有一次协调器所在的节点因为磁盘满了导致进程崩溃,整个训练任务停了两个小时才发现。
解法是给协调器加一个热备节点。主协调器和备协调器之间通过一个轻量级的共识协议保持状态同步(我们用的是基于Raft的简化实现)。主协调器定期将状态快照发送给备协调器,备协调器在检测到主协调器心跳超时后自动接管。切换过程中,训练侧和推理侧的心跳会短暂丢失,但由于我们有两级故障判定机制,不会触发误判。切换完成后,新的主协调器会从最近的checkpoint恢复状态,继续协调工作。
6. 性能开销的实测数据与优化建议
6.1 checkpoint保存对训练吞吐的影响
我们做了一组对比实验,测量checkpoint保存对训练吞吐的影响。实验设置:模型规模7B,训练侧4个节点(每个节点8张A100),推理侧2个节点(每个节点8张A100),rollout batch size 64,训练batch size 32。
在不保存checkpoint的情况下,训练吞吐大约是每秒1.2个step。在每100个step保存一次checkpoint的情况下,训练吞吐下降到每秒1.1个step,下降了约8%。这个开销主要来自参数冻结期间的等待时间。如果把checkpoint间隔拉长到每500个step,吞吐下降只有2%左右。
我们的建议是:checkpoint间隔不要设得太短,除非你的训练任务非常不稳定。一般来说,每200到500个step保存一次是一个比较合理的范围。如果训练任务的单次运行时间很长(比如超过24小时),可以适当缩短间隔,但不要短于100个step,否则协调开销会显著影响训练效率。
6.2 参数同步频率与恢复精度的权衡
参数同步频率直接影响恢复精度。同步越频繁,恢复时需要回滚的步数越少,恢复后的训练越接近故障前的状态。但同步越频繁,协调开销越大,训练吞吐越低。
我们测试了几组不同的同步频率:每1个step同步一次、每3个step同步一次、每5个step同步一次。结果显示,每1个step同步时,恢复后的loss曲线和故障前几乎无缝衔接,但训练吞吐下降了约15%。每5个step同步时,训练吞吐只下降3%,但恢复后需要大约20个step才能回到故障前的loss水平。
综合考虑,我们最终选择了每3个step同步一次。这个频率下,训练吞吐下降约7%,恢复后大约5到10个step就能回到正常水平。当然,这个选择取决于具体的训练任务和硬件配置,没有一个通用的最优值。建议在实际部署前做一组小规模的对比实验,找到适合自己场景的平衡点。
6.3 存储IO的瓶颈与缓解
checkpoint的存储IO是另一个容易被忽视的瓶颈。我们的模型参数加上优化器状态大约40GB,写入分布式文件系统的速度大约是200MB/s,单次checkpoint的写入时间大约200秒。如果checkpoint间隔是200个step,每个step耗时约0.8秒,那么checkpoint写入时间占训练总时间的比例大约是200/(200*0.8+200)≈55%。这个比例太高了,严重影响了训练效率。
缓解方案有几个:一是使用更快的存储介质,比如NVMe SSD阵列,我们把写入速度提升到了1GB/s,写入时间缩短到40秒。二是采用增量checkpoint,只保存发生变化的那部分参数,但增量checkpoint的实现复杂度较高,而且恢复时需要重放历史变更,我们评估后认为不适合我们的场景。三是异步checkpoint,在后台线程中写入checkpoint,不阻塞训练主流程,但需要额外的内存来保存checkpoint的快照,对内存压力较大。
我们最终采用的是"快速存储+异步写入"的组合方案。checkpoint先写入本地NVMe SSD(速度快),然后由后台线程异步同步到分布式文件系统(保证可靠性)。这样训练主流程只需要等待本地写入完成(大约40秒),后台同步不影响训练。这个方案在保证可靠性的同时,把checkpoint对训练吞吐的影响降到了最低。
6.4 不同规模下的参数调优参考
最后给出一组我们在不同模型规模下的参数配置参考,供读者根据自己的场景调整:
| 模型规模 | checkpoint间隔 | 同步频率 | 心跳间隔 | 观察窗口 | 冻结超时 |
|---|---|---|---|---|---|
| 1B-3B | 500 step | 5 step | 10s | 30s | 15s |
| 7B-13B | 300 step | 3 step | 15s | 60s | 30s |
| 30B-70B | 200 step | 2 step | 20s | 90s | 45s |
| 100B+ | 100 step | 1 step | 30s | 120s | 60s |
这张表里的数值不是绝对的,需要根据实际的网络环境、存储性能和训练稳定性来调整。比如如果你的集群网络抖动频繁,观察窗口应该适当加长;如果你的存储写入速度很快,checkpoint间隔可以适当缩短。
我个人在实际操作中的体会是,Checkpoint Engine的接入不是一个"一次性配置好就不用管"的事情。随着训练任务的变化(模型规模、数据分布、集群负载),最优的参数配置也会变化。建议在训练过程中持续监控checkpoint的成功率、恢复时间和训练吞吐,定期回顾和调整配置。另外,故障恢复流程一定要定期演练,不要等到真正出故障的时候才发现恢复脚本有bug。我们现在的做法是每周做一次模拟故障演练,随机kill掉一个推理节点,验证恢复流程是否正常。这个习惯帮我们提前发现了好几个潜在的恢复问题。