
说个我自己的经历。最开始用 STABLE_BASELINE3 跑 PPO 的时候环境是 CartPole几十秒就能出一版结果所以从来没想过模型保存这件事。后来换到一个机械臂仿真环境一次训练动辄两三个小时某天晚上笔记本合盖休眠训练进程直接没了。那一刻我才发现模型保存、读取、再训练这三件事情才是实战里最容易忽略、也最基础的基本功。这篇文章就把 SB3 里和模型管理相关的操作完整梳理一遍保存什么、存在哪、怎么读回来、如何拿着老模型接着训。无论你是刚入门强化学习、正在跑第一个 SB3 项目还是已经训练过一两个模型但没研究过 save/load 机制都可以照着代码一步步过一遍。1. 为什么说模型管理是强化学习实战的必修课1.1 训练过程太脆弱不保存就是在赌运气很多人一开始在 Gym 的经典环境上跑实验CartPole 几千步就收敛训练时间比写代码还短。但真实项目里完全不是这么回事机器人控制、推荐系统、交通调度、工业优化这类场景观察空间动辄几十维单次交互又要跑仿真甚至真机训练步数轻松上百万。这时候一次训练的耗时是以小时甚至天来计算的。强化学习训练又比普通监督学习更容易被打断。普通深度学习训练哪怕进程崩了模型参数基本已经收敛损失曲线也能看到大概趋势。强化学习不一样策略是一步一步和环境交互试出来的任何一个阶段的模型都是中间产物一旦没有存盘进程崩溃、机器重启、云端资源被回收之前所有时间等于直接归零。我踩过最痛的一次是在远程服务器上跑一个连续控制任务训练到第 22 万步断了 SSH训练进程没挂到 nohup 下面登录上去一看进程没了而代码里没有配置任何自动保存。那次之后我把模型管理直接列进了模板代码的第一行注释里训练可以断模型必须留。所以模型保存不是可选优化项而是训练流程里的保底机制。它解决的不只是崩溃恢复这一个问题后面整个实验管理、调试、对比、部署都建立在它之上。1.2 模型保存不只是“存档”还支撑着三个关键场景断点续训是最直观的场景。训练到一半因为各种原因中止想要从最近的位置接着跑就必须有中间模型存档。这里有个细节存档最好定期落盘而不只保存在内存变量里。很多人手动model.save()之后没有用回调定期保存等到进程真的崩了才发现最后一次保存已经是很久之前的状态相当于中间一大段训练全部浪费。第二个场景是超参数对比和策略微调。比如你已经有一个在 Pendulum-v1 上训练得不错的 PPO 模型现在想测试不同学习率、不同 clip 范围、或者不同奖励函数改动之后的效果。如果你每次都从随机初始化开始训练收敛时间和方差都会干扰你的判断。但如果你从一个已保存的预训练模型出发只让它继续训练一小段就能更清晰地看到超参改动带来的差异。这也是再训练最常见的使用方式。第三个场景是部署推理。训练和推理往往是两个阶段、两个环境。训练时模型在内存里环境在仿真器里部署时模型可能要被封装进服务、打包到边缘设备甚至跑在完全不同的操作系统里。这时候你需要的是一个独立的模型文件而不是训练脚本里的 model 变量。把训练好的模型save()出来在部署侧用load()读进去直接就省掉了把网络结构抄一遍并手动复制权重这种麻烦事。2. 保存模型先搞懂 SB3 到底存了什么2.1 save() 一行代码背后权重、优化器、空间定义全在里面最简单的保存方式是这样的from stable_baselines3 import PPO model PPO(MlpPolicy, CartPole-v1, verbose1) model.learn(total_timesteps10000) model.save(ppo_cartpole)运行完之后当前目录会多出一个ppo_cartpole.zip文件。这个名字很容易让人以为 SB3 只是把神经网络权重打包了一下实际上里面装的东西比权重多得多。模型文件里至少包含了这几个部分policy 网络的参数这是最核心的推理内容optimizer 的状态包括 SGD/Adam 中的动量、方差等参数。你有没有想过为什么load()之后还能继续训练得很好就是因为优化器状态也存下来了续训不会出现模型已经收敛但优化器还在用小步长乱跳的割裂感learning rate schedule 相关的调度器状态observation_space 和 action_space 的定义训练元信息比如 seed、num_timesteps、policy_class、policy_kwargs 等。保存一行代码恢复也是一行代码。但明白里面有什么对你理解为什么续训前要小心处理学习率很重要。后面第 4 部分会详细展开。顺带提醒一句官方文档的习惯是保存路径不带后缀让 SB3 自动补.zip。你自己写上ppo_cartpole.zip一般也能正常加载只是没必要。保持一种统一写法团队协作时少很多歧义。2.2 自动保存实战CheckpointCallback 与 EvalCallback 配合使用手动save()只能保存当前时刻的模型但训练过程中你不会一直在旁边盯着现在该存了。SB3 提供了CheckpointCallback可以按固定时间步自动保存from stable_baselines3.common.callbacks import CheckpointCallback checkpoint CheckpointCallback( save_freq20000, save_path./checkpoints/, name_prefixppo_cartpole, ) model.learn(total_timesteps200000, callbackcheckpoint)这里的save_freq20000表示每训练 20000 个环境交互步保存一次不是每 20000 个 episode。如果你希望按 episode 来保存就得自定义回调在里面数 done 的次数再触发保存。另一个更智能的组件是EvalCallback。它会按固定间隔用一组长度的评估 episode 测试当前模型并把历史最高评估分数的模型单独保存成best_model.zipfrom stable_baselines3.common.callbacks import EvalCallback eval_callback EvalCallback( eval_enveval_env, best_model_save_path./best/, log_path./eval_logs/, eval_freq5000, n_eval_episodes10, deterministicTrue, ) model.learn(total_timesteps200000, callback[checkpoint, eval_callback])EvalCallback的价值在于它可以同时解决两个痛点一个是最新的不一定是最好的强化学习训练后期策略可能波动甚至退化有了 best_model 你可以随时回溯到历史最优另一个是它能实时记录评估指标到 tensorboard让你不用每次都手动跑一段测试脚本就能看到平均 reward 的走势。我自己使用时的习惯是CheckpointCallback保存进度节点EvalCallback保存最优节点。前者用于断点续训后者用于找到最终部署的模型。两者各司其职不要混成一个。2.3 命名规范和路径处理避免文件管理翻车模型文件一旦变多最头疼的不是训练而是这个文件是谁。我以前就吃过亏一个目录下堆了model.zip、model(1).zip、final_model.zip等到想回滚到某个实验版本根本分不清哪个对应哪次实验。后来我统一了命名格式推荐给你参考{algorithm}_{env}_{total_steps}_{timestamp}.zip例如ppo_pendulum_100000_steps_202501171430.zip。如果你用CheckpointCallbackname_prefix里直接带上算法和环境的组合文件名自动拼上步数这样至少能一眼看出是哪个算法、哪个环境、训练到第几步。路径方面最容易踩的坑是相对路径。训练脚本从项目根目录启动时./checkpoints/指向根目录下的文件夹但如果你换到子目录启动或者写进 systemd/tmux 服务脚本里相对路径就会漂移导致保存的位置和你预想的不一样。我现在所有训练脚本里都会统一用pathlib生成绝对路径from pathlib import Path BASE_DIR Path(__file__).resolve().parent MODEL_DIR BASE_DIR / models MODEL_DIR.mkdir(exist_okTrue)然后在所有 save/load 处都用MODEL_DIR / xxx至少能保证路径不随启动目录变化。3. 读取模型load 的没那么简单这些坑你得知道3.1 PPO.load() 关键参数逐个拆解读取模型一行代码model PPO.load(ppo_cartpole, envenv, devicecpu)这里我故意把常用参数都写上方便逐个说明。path很直接填模型文件路径就行。注意如果是相对路径和保存时一样启动目录变了就会出问题。env是最容易让人疑惑的参数。其实load()时可以不传 env模型照样能加载进来predict()也能跑。但如果你加载之后想继续learn()就必须先绑定环境。传env就相当于自动调用了set_env()。如果你 load 时没传后面也还有补救机会model.set_env(env) model.learn(total_timesteps10000)device用来指定模型加载到哪个设备。SB3 默认会沿用保存时的 device所以在 GPU 上训练的模型加载到只有 CPU 的环境时会报 device 不匹配或者自动回退到 CPU。为了明确我一般在部署场景里都会显式传devicecpu。最后是custom_objects。如果你训练时用了自定义 policy、自定义 feature extractor或者自定义的 RL 模块load 的时候 SB3 需要知道这些自定义类的实现。用字典形式传入model PPO.load( ppo_custom, custom_objects{policy_class: MyCustomPolicy}, )这个参数平时用不上但一旦用上就是救命级别的。否则你会看到类似ValueError: Cannot find module ...的报错明明 zip 文件就在那里却怎么都加载不进来。3.2 预测时动作“不动了”deterministic 与随机策略的区别模型加载成功后很多人会直接写action, _ model.predict(obs)然后发说模型表现很好但动作好像统一起不来 —— 其实这往往不是 bug而是predict()默认行为和随机策略的采样逻辑导致的。SB3 的predict()有个参数deterministic默认值是False。对于 PPO、A2C 这类带随机策略网络的算法来说deterministicFalse意味着每一步都会从动作分布里采样行为会带有明显的随机性而deterministicTrue则是直接取概率分布的均值作为动作行为更稳定。我在实际评估里几乎永远用deterministicTrueaction, _ model.predict(obs, deterministicTrue)为什么因为强化学习评估要的是这个策略在当前状态下的稳定表现。如果用随机采样同一模型在同一个起点跑十次结果可能相差不小这会给不同模型之间的横向对比带来噪声。测试时想观察模型的探索行为可以把deterministic设为False再跑一两次但做评估和最终部署决策时请统一用确定性模式。还有一个容易误解的点对于 DQN 这类 off-policy 算法保存模型里并不会有训练时的 epsilon 探索参数在推理阶段生效。predict()在推理时本身就等同于贪心策略所以你不需要太过担心 explore/exploit 的问题。3.3 环境不匹配、算法类不匹配两个高频报错解析我见过最多、也最容易排查的就是环境不匹配问题。训练时用的是Box(4,)的观测空间加载后却绑定了一个观测空间是Box(8,)的环境或者动作空间从离散变成连续。这种情况下 SB3 不会非常礼貌地提示你的环境换错了它通常会在predict()或learn()里报一个维度错误甚至直接抛assert失败。排查思路很直接先看加载时报错信息里的 space 描述再去训练脚本里核对环境 ID。最保险的做法是训练脚本和续训脚本共用同一个环境注册入口不要手动改observation_space。第二种常见报错是算法类不匹配。比如你训练时用的是DQN加载时却写成PPO.load(...)。SB3 的模型文件里记录了训练时的算法类名加载时它会检查当前调用的类是否匹配不匹配就会报错。这类错误的名字可能带class、expected这类字样看到后第一反应就是这次调用和保存时的算法类不一致。另外policy_kwargs不一致也会导致类似问题。比如训练时你定义了net_arch[256, 256]加载时没有传这个配置并且网络结构对不上同样会报错。解决办法就是保持并集PPO.load(ppo_custom, envenv, policy_kwargsdict(net_arch[256, 256]))或直接自定义 policy 类。4. 再训练继续学习不是无脑 load learn4.1 续训前必须检查的三件事学习率、经验池、探索设置现在到了文章的重头戏再训练。最直觉的续训代码是model PPO.load(ppo_cartpole, envenv) model.learn(total_timesteps50000)代码没错能跑但直接这么用经常会得到越训越差的结果。我总结下来续训前必须检查三件事。第一学习率。保存模型时SB3 把 learning rate schedule 也存进去了。如果你训练时用的是线性衰减的 schedule保存时 schedule 已经走了一段生命周期。但是问题在于SB3 的load()在恢复时并没有办法知道你这个再训练目标总步数是多少所以它恢复的调度器状态可能和你续训总长度对不上。最典型的后果就是你以为它在用小学习率微调实际却用了一个偏大的学习率重新起步策略直接被冲散。解法是在load()时显式指定一个新的学习率model PPO.load(ppo_cartpole, envenv, learning_rate1e-4)我的习惯是续训阶段的学习率设为初始训练学习率的 0.3~0.5 倍。原先 3e-4续训就 1e-4这样既能继续优化又不至于把已有策略完全打乱。第二经验池。这一点对 PPO 这类 on-policy 算法影响不大因为每次更新后 rollout buffer 本来就会清空。但如果你在续训 DQN、SAC、TD3 这类 off-policy 算法问题就大了。SB3 默认保存模型时不会保存 replay buffer也就是说你辛辛苦苦攒了十几万条经验样本load()回来之后 buffer 是从零开始的。这会让续训前段退化成重新探索严重拖慢收敛速度。第三探索设置。DQN 在训练前期依赖 epsilon 探索epsilon 会随时间衰减。如果你续训时没有显式设置exploration_fraction加载回来的模型很可能已经处于探索几乎停止的状态。模型如果还没完全收敛那它后续很难跳出局部最优。DDPG/TD3 这类算法也有类似问题加载后续训时注意 action noise 的sigma是否需要重置。PPO 相对好一些探索基本体现在策略分布的方差上续训时只要学习率不过大方差不会骤然坍缩。4.2 完整实操从一个老模型继续训练并保留监控指标下面给一个完整的续训示例环境用 Pendulum-v1第一阶段训 2 万步第二阶段加载后训 5 万步import numpy as np from stable_baselines3 import PPO from stable_baselines3.common.env_util import make_vec_env from stable_baselines3.common.callbacks import CheckpointCallback, EvalCallback # 阶段一初始训练并保存 env make_vec_env(Pendulum-v1, n_envs1) model PPO(MlpPolicy, env, learning_rate3e-4, verbose1) model.learn(total_timesteps20000) model.save(ppo_pendulum_stage1) # 阶段二加载老模型降低学习率续训 env make_vec_env(Pendulum-v1, n_envs1) model PPO.load(ppo_pendulum_stage1, envenv, learning_rate1e-4) checkpoint CheckpointCallback( save_freq10000, save_path./checkpoints/, name_prefixppo_pendulum_retrain, ) eval_callback EvalCallback( eval_envmake_vec_env(Pendulum-v1, n_envs1), best_model_save_path./best/, log_path./eval_logs/, eval_freq5000, n_eval_episodes10, deterministicTrue, ) model.learn( total_timesteps50000, callback[checkpoint, eval_callback], reset_num_timestepsFalse, ) model.save(ppo_pendulum_stage2)注意最后这行model.save(...)在learn()之后再执行一次一定别漏。EvalCallback只会把评估最高的保存为best_model.zip但不会自动保存最终模型的完整状态。我自己就曾经靠best_model回滚却发现没有保存最终训练完那一刻的模型后面想复现当时的日志状态很麻烦。关于reset_num_timestepsFalse这个参数的作用是告诉 SB3 不要重置训练步数计数。如果保持默认第二次learn()会从 0 开始计数tensorboard 里会出现两条重合的曲线很难判断续训到底提升了多少。设成 False 后曲线时间轴连续效果一目了然。如果你希望每个阶段干净地分开观察那么设成 True 也可以只是指标对比时要自己注意阶段边界。续训时一定要把EvalCallback开着。强化学习训练不是只往上涨不往下跌的尤其是从旧模型继续训练策略可能先快速下降一阵子再恢复。如果没有 eval 曲线单靠训练日志里的 loss 根本判断不了策略是不是真的在变好。4.3 更稳妥的存档方案手动保留 replay buffer刚才说到 SB3 默认不保存 replay buffer那么对 off-policy 算法来说怎么把 buffer 也一起保存下来最简单的方式是直接 pickleimport pickle # 模型保存之外把 replay buffer 单独存一份 with open(dqn_replay_buffer.pkl, wb) as f: pickle.dump(model.replay_buffer, f) # 续训时加载 model DQN.load(dqn_cartpole, envenv) with open(dqn_replay_buffer.pkl, rb) as f: model.replay_buffer pickle.load(f)这个方案在 SB3 2.x 中实测可行前提是保存和加载两端的环境、buffer 类结构保持一致。如果你跨 Python 版本或跨机器使用时出现 pickle 兼容问题那就换一种更规范的方案把 buffer 里的 observation、action、reward、done 等数组分别用np.save保存下来加载后重新组装一个 ReplayBuffer 实例再填充。代码会多一点但更稳也更容易排查。如果你经常做 DQN 续训建议把你自定义的保存模型 保存 buffer功能封装成一个回调或者写成两个工具函数save_checkpoint_with_buffer(model, path)和load_checkpoint_with_buffer(model_class, path, env)。否则每次续训都手动pickle.dump一次很容易漏。这里有一个容易忽略的细节DQN 在预测阶段并不需要 replay buffer但续训阶段没有 replay buffer 就是灾难。所以如果你的项目里保存模型和保存 buffer是两个环节请确保续训脚本把这两件事都做了否则你会看到一个很反直觉的现象eval分数正常甚至不错但一learn()loss 和 reward 立刻大幅波动。5. 高频问题排查清单与避坑记录5.1 问题速查表我把实际使用中见过的问题整理成一张速查表方便你在遇到问题时直接对照现象根本原因解决办法load()后learn()报 env 相关错误加载时没传 env续训前没绑定环境load(..., envenv)或model.set_env(env)续训后 reward 不升反降学习率设置不当 / replay buffer 丢失 / 探索不足降低 learning_rate恢复 buffer必要时重设探索参数predict()输出动作长时间不变随机策略下用了deterministicTrue且均值动作稳定评估用 True需要探索行为时用 FalseDQN 加载报 class mismatch 类错误用错了算法类确认保存时用的算法用相同类加载加载后predict()维度报错观测/动作空间不匹配保证训练与加载环境一致核对 space模型存档太大/太小完整存档包含优化器状态replay buffer 未保存训练存档用完整 zip部署只导 policy 参数tensorboard 训练曲线从 0 开始续训时没设reset_num_timestepsFalse显式传参数保持时间轴连续自定义网络加载失败自定义 policy 类找不到使用custom_objects传入类或类路径补充一个容易被忽略的如果你只是想把一个模型的参数复制给另一个同结构的模型不想经过文件 IO可以直接用get_parameters()和set_parameters()。这在超参对比实验中很好用能省掉频繁读写磁盘的开销params model1.get_parameters() model2.set_parameters(params)但注意set_parameters只是替换网络参数不会复制优化器状态。如果需要完整状态迁移还是用保存/加载 zip 文件更合适。5.2 几个值得长期坚持的实操习惯最后分享几个我自己一直在坚持的实操习惯希望能帮你少走弯路。训练脚本的入口处统一配置好目录。MODEL_DIR、LOG_DIR、CHECKPOINT_DIR都用绝对路径生成不要散落在代码里用相对路径。训练一开始就挂上CheckpointCallback和EvalCallback哪怕只是跑 5 分钟的小实验也挂上。习惯成自然之后你永远不会再遇到从头跑一遍的尴尬。给每个实验留一份配置记录。最简单的做法是训练结束时把超参数、环境 ID、Git commit、数据集版本写进一个 JSON 文件放在模型同目录下。这样后面看到ppo_pendulum_stage2.zip还能知道它是用什么环境、什么学习率、跑了多少步得到的。模型文件按周期清理。CheckpointCallback默认会一直存到训练结束步数多了之后文件会越来越多。如果你只在续训时有意义那么只保留最近 3~5 个 checkpoint 加上best_model.zip就足够。最后部署阶段如果不需要继续训练只关注推理性能可以考虑只导出 policy 参数不导出优化器状态。这样模型文件更小、加载更快。但如果是为了续训和实验管理一定要保存完整 zip。我自己跑实验到现在最狼狈的一次不是训练算法没收敛而是跑了一整晚的 PPO 因为中途忘记配置 checkpoint第二天起来发现进程崩了只剩一个初始模型。从那以后我的代码模板里CheckpointCallback永远是第一个写进去的组件。模型管理这件事前期多花五分钟后期真的能帮你省下几个小时的重训时间。希望这篇笔记能让你少走一点类似的弯路。