强化学习模型管理:SB3保存、加载与再训练实战指南

发布时间:2026/9/15 15:39:51
强化学习模型管理:SB3保存、加载与再训练实战指南 训练一个强化学习模型最怕的是什么不是收敛慢是辛辛苦苦跑了一晚上的实验第二天发现没保存全白干了。我刚上手 Stable Baselines3SB3那会儿踩的第一个大坑就是这个——以为 model.learn() 结束就等于万事大吉结果一次环境崩溃几天小时的训练直接归零。后来花了不少精力把模型保存、模型读取、再训练这条链路彻底摸清才发现这才是强化学习实战里真正决定项目能不能落地、能不能长期迭代的关键。这篇就把 SB3 下完整的模型生命周期管理方法论、踩过的坑和可复现代码整理出来适合刚入门强化学习的同学也适合正在做断点续训、策略微调、模型上线的老手参考。1. 先把三件事串起来保存、读取、再训练在实战中的位置1.1 为什么模型保存是训练断点而不是作业存档很多人在学强化学习时用的都是 Gym 环境里小规模的 CartPole、Pendulum几分钟就能训练好所以对模型保存这件事并不敏感觉得反正随时能重新训。一旦进入真实项目比如机械臂抓取、组合优化求解、推荐系统策略环境的交互成本会暴增。一个 PPO 在复杂环境里跑几百万步花掉的可不只是 GPU 上的几小时还有仿真器的计算资源和大量的数据采集时间。这时候训练中断、参数调错、奖励函数改版都是家常便饭。模型保存就成了那个让你不至于从零开始的训练断点和普通作业写完随手保存一个概念完全不同。另外强化学习的训练环境和部署环境往往是分离的。训练机上你可以用 128 个并行环境跑一整天但部署到边缘设备或生产环境时不可能带着训练脚本跑。你需要的是一个已经收敛好的策略文件加载进来直接做推理。没有模型读取这步训练得再好也落不了地。1.2 一套贯穿项目始终的模型生命周期管理我理解中的模型生命周期管理包括三件事保存、读取、再训练。保存不只是把权重写进磁盘要保证优化器状态、环境配置、归一化统计量、经验回放池这些附带的上下文都能按需存档读取也不是机械地 load 一个文件要搞清楚哪些信息在加载时自动恢复、哪些必须手动传参再训练则分成两类完全不同的场景一类是接续之前的训练进度继续跑另一类是拿到一个已经训练好的模型做微调或迁移。很多人只会调用 API但对这三件事的边界和坑一无所知导致项目写到一半才发现问题。这篇的内容全部基于 Stable Baselines3 2.x 版本默认配合 gymnasium 环境库使用。如果你还在用 SB3 1.x 和旧版 gym部分 API 和数据结构有差异阅读时需要留意。2. 模型保存四种存档方式和它们背后的坑2.1 最基础的 model.save() 与 zip 的庐山真面目先看最常用的保存姿势from stable_baselines3 import PPO import gymnasium as gym env gym.make(CartPole-v1) model PPO(MlpPolicy, env, verbose1) model.learn(total_timesteps50000) model.save(ppo_cartpole)这段代码执行后会在当前目录生成一个ppo_cartpole.zip文件。你可以把它看成一个模型快照包里面不仅包含 PyTorch 的 policy 权重还包含算法超参数、观察空间、动作空间、策略类信息等元数据。这也是很多初学者容易忽略的地方——SB3 的 save() 不是简单地把权重序列化它是一个打包了完整上下文的存档机制。用解压工具打开这个 zip 会看到两类文件一个是保存超参数和空间信息的 data 文件另一个是 policy.pth 格式的 PyTorch 权重。正因为 SB3 把超参数也存进去了PPO.load()才能在新环境里恢复出一个和原模型配置完全一致的对象。注意model.save(ppo_cartpole)和不带扩展名或手动加.zip效果一样。但如果你写model.save(ppo_cartpole.pth)文件名会带.pth后缀这并不会让保存格式变成 pth它仍然是一个 zip 包。保存格式由save_format参数控制跟扩展名没有直接关系。2.2 CheckpointCallback 自动存档训练崩了也不慌手动 save 只能覆盖一个时间点的快照。实际训练中策略是动态变化的训练到后期往往还会出现性能回退。只保存最后一个模型很可能最后保存到的恰好是过拟合或奖励震荡后的烂模型。我建议从一开始就给 learn() 挂上 CheckpointCallbackfrom stable_baselines3.common.callbacks import CheckpointCallback checkpoint CheckpointCallback( save_freq10000, save_path./checkpoints/, name_prefixppo_cartpole, save_vecnormalizeTrue, ) model.learn(total_timesteps200000, callbackcheckpoint)这样每 10000 步就会自动生成一个带时间步标记的存档比如ppo_cartpole_10000_steps.zip、ppo_cartpole_20000_steps.zip。训练中途崩溃损失最多只占一个 save_freq 的间隔训练结束后想回退到某个中间状态也可以直接挑一个时间拾取点加载非常实用。save_vecnormalizeTrue这个参数值得单独说一下。如果训练脚本里用了 VecNormalize 包装环境后面专门展开训练过程中会不断更新观测值和奖励的均值方差统计。用带这个参数的 CheckpointCallback 保存时会额外生成一个对应的vecnormalize.pkl文件把环境统计量和模型一起存档加载时配套使用才能保证环境输入分布一致。2.3 别忘了 VecNormalize环境统计信息要单独保存这是 SB3 实战里模型能保存但加载后跑飞了的头号原因。VecNormalize 是 SB3 官方用于对观测和奖励做归一化/标准化的包装器很多场景下对收敛速度的提升非常明显。但它的均值和方差是训练中在线更新的不会自动写进 model.save() 的 zip 包里。如果只保存模型、不保存 VecNormalize下次加载模型接入一个全新的 VecNormalize统计量是初始值模型看到的观测分布和训练时完全不同表现必然崩坏。正确的保存方式是把两者绑定保存from stable_baselines3.common.vec_env import DummyVecEnv, VecNormalize venv DummyVecEnv([lambda: gym.make(CartPole-v1)]) venv VecNormalize(venv, norm_obsTrue, norm_rewardTrue, clip_obs10.0) model PPO(MlpPolicy, venv, verbose1) model.learn(total_timesteps100000) model.save(ppo_cartpole_vec) venv.save(vecnormalize_cartpole.pkl)加载时一定不要忘了恢复它venv DummyVecEnv([lambda: gym.make(CartPole-v1)]) venv VecNormalize.load(vecnormalize_cartpole.pkl, venv) # 推理时要把 training 置为 False否则统计量还会继续更新 venv.training False venv.norm_reward False model PPO.load(ppo_cartpole_vec, envvenv)如果你是用来继续训练记得把venv.training重新设回 True如果只是加载模型做评估或部署就保持 False并关闭norm_reward避免评估时环境还在偷偷更新统计量。2.4 用 save_format 控制存档格式SB3 的 save() 支持save_format参数可选zip默认和pth。zip 格式是最推荐的主力格式因为它同时保存了算法类信息、超参数和策略权重PPO.load()可以直接通过这个文件恢复完整模型。pth 格式则只保存 PyTorch 的 state_dict适合你只是想导出权重做自定义推理、或者要把权重迁移到别的框架时的场景model.save(ppo_cartpole_only_weights.pth, save_formatpth)但 pth 格式没有超参数和空间信息无法用 PPO.load() 直接完整恢复。如果你只是想要一个权重中转站pth 够用如果是做实验管理、断点续训请务必使用默认的 zip 格式。我在实践中会把两种格式区分开zip 格式用于训练存档pth 格式用于最终部署导出。这样既保证了实验可回滚也方便下游工程取用纯权重。2.5 经验回放缓冲Off-policy 算法的最值钱资产如果你用的是 DQN、SAC、TD3 这类 off-policy 算法除了策略权重还有一个非常重要的资产叫 replay buffer经验回放池。它保存了 agent 与环境交互过的所有转移样本是 off-policy 算法的核心数据。训练中断后只加载模型、不加载 replay buffer模型确实能接着跑但回放池是空的必须重新探索收集一批数据才能进入正常学习节奏前期会有一段明显的性能滑坡。SB3 为 off-policy 算法专门提供了接口from stable_baselines3 import SAC model SAC(MlpPolicy, env, verbose1) model.learn(total_timesteps50000) model.save(sac_pendulum) model.save_replay_buffer(sac_replay_buffer) # 恢复时 model SAC.load(sac_pendulum, envenv) model.load_replay_buffer(sac_replay_buffer) model.learn(total_timesteps50000)注意save_replay_buffer只适用于 off-policy 算法。PPO、A2C 这类 on-policy 算法没有跨训练阶段生效的经验回放池它们每轮 rollout 后立即更新再丢弃所以不需要也无法单独保存 replay buffer。如果你强行调用会得到 NotImplementedError。3. 模型读取load() 不是简单的解压文件3.1 一份模型两种打开方式评估 vs 继续训练先说结论PPO.load(ppo_cartpole)得到的模型对象没有绑定任何环境。它内部已经保存了 observation_space 和 action_space可以直接用来做推理预测但没法直接调用learn()继续训练——继续训练必须传入环境。model PPO.load(ppo_cartpole) # 评估/推理 obs, info env.reset() for _ in range(1000): action, _ model.predict(obs, deterministicTrue) obs, reward, terminated, truncated, info env.step(action) if terminated or truncated: obs, info env.reset()这里有两个细节。第一predict(obs, deterministicTrue)表示选择确定性动作即策略网络输出的均值或 argmax如果设成 False则会从动作分布中采样适合训练或探索阶段。二是 predict 返回的是一个二元组(action, state)第二个元素是 RNN/GRU 等循环策略使用的隐藏状态普通 MLP 策略直接忽略即可。如果你要接着训练需要在 load 时把环境传进去env gym.make(CartPole-v1) model PPO.load(ppo_cartpole, envenv) model.learn(total_timesteps30000)3.2 自定义网络加载policy_kwargs 必须原样对齐如果你的策略网络是自定义的加载时有一个很隐蔽的坑。SB3 保存 zip 时会把policy_kwargs原样写进 data 文件加载时本应自动恢复。但问题出在如果自定义类是在脚本里临时定义的函数内部创建的匿名类或者类的引用路径在加载时不可达SB3 就找不到这个类无法正确恢复网络结构。正确做法是一类是把自定义网络定义在模块顶层保证加载时 Python 能 import 到另一类是加载时通过custom_objects参数手动指定from stable_baselines3.common.torch_layers import BaseFeaturesExtractor import torch.nn as nn class MyFeatureExtractor(BaseFeaturesExtractor): def __init__(self, observation_space, features_dim256): super().__init__(observation_space, features_dim) self.net nn.Sequential( nn.Linear(observation_space.shape[0], 256), nn.ReLU(), nn.Linear(256, features_dim), ) def forward(self, observations): return self.net(observations) policy_kwargs dict(features_extractor_classMyFeatureExtractor) model PPO(MlpPolicy, env, policy_kwargspolicy_kwargs) model.learn(20000) model.save(ppo_custom_extractor) # 在另一个脚本中加载 model PPO.load( ppo_custom_extractor, envenv, custom_objects{policy_kwargs: policy_kwargs}, )自定义网络时另一个常见问题是features_dim不匹配。保存前网络输出维度是 256加载时如果自定义对象里写成 128策略头部维度对不上会直接报维度错误。我建议把 policy_kwargs 集中写在一个配置字典里保证训练和加载共用同一份配置。3.3 模型文件里到底存了什么为什么经常报错遇到加载报错时第一反应不要去猜直接拆开模型 zip 看看里面存了什么。在命令行执行unzip -l ppo_cartpole.zip输出里能看到 data 文件和 policy.pth 文件。data 文件里保存了算法类名、policy_class、hyperparameters、observation_space、action_space 等。如果你发现加载报错提示空间维度对不上最可能的原因是训练环境和加载环境的空间定义不一致。比如训练用的是Discrete(2)加载时接入了一个Discrete(3)的环境模型当然翻车。还有一类报错发生在跨 SB3 版本加载时。同一个模型文件在 SB3 1.8 和 SB3 2.2 之间不一定能无缝加载因为内部序列化格式有过调整。如果项目跨越了多个版本长期迭代建议固定 SB3 版本或者升级后用旧模型重新测试一轮再决定是否复用。4. 再训练实战断点续训和策略微调的区别与实现4.1 断点续训的标准流程再训练的第一种典型场景是断点续训——就是上次训练因为中断或资源限制没跑完这次接续之前的进度继续学习。操作本身很简单env gym.make(CartPole-v1) # 第一步加载模型同时绑定环境 model PPO.load(ppo_cartpole_50000_steps, envenv) # 第二步继续训练这里的 total_timesteps 是增量不是累计目标 model.learn(total_timesteps100000) # 第三步重新保存 model.save(ppo_cartpole_150000_steps)需要注意total_timesteps始终表示这一次调用 learn 要跑多少步而不是从 0 到某个总步数。SB3 官方的设计是你每次调用 learn 传入的是本次新增的步数。如果你理解成总目标会多跑很多不必要的轮次。4.2 学习率与优化器状态的残酷真相很多人以为model.load()之后继续learn()会恢复优化器的动量等状态因为保存时存了优化器快照。但 SB3 的实际行为是每次调用 learn() 时都会重新创建优化器优化器的动量不会从存档中恢复。也就是说加载后继续训练只有网络参数是继承的优化器状态是从零开始的。这个设计影响很大。如果你原来用的是衰减学习率调度比如从 3e-4 衰减到 1e-5那加载后继续训练学习率调度也会重新开始。如果你希望在第二阶段用一个相对更小的学习率做收敛直接在加载后手动覆盖model PPO.load(ppo_cartpole, envenv) model.learning_rate 1e-4 model.learn(total_timesteps50000)如果你的学习率是通过 schedule 函数传入的比如learning_ratelambda progress_remaining: 3e-4 * progress_remaining那 leran 时会根据这个函数重新从头开始衰减。想要第二阶段的衰减幅度小一些可以用另一个 lambda比如lambda progress_remaining: 1e-4 * (0.5 0.5 * progress_remaining)让初始学习率减半衰减曲线也平缓一些。这块没有标准答案但第二阶段必降学习率是强化学习实战中的一条普适经验。4.3 从续训到微调奖励函数改了怎么办第二类再训练场景是策略微调。最常见的起因是奖励函数改版了——要么是原来的奖励稀疏导致学不动要么是任务目标局部调整。这种情况下你手里有一个在旧奖励下已经学到不少策略的模型直接丢掉重新训练太浪费。正确做法是加载旧模型接入新环境调低学习率再训练少量步数让策略平滑迁移到新奖励分布下。需要注意如果奖励函数改动较大旧策略产生的数据分布和新目标可能差别很大微调初期会出现性能短暂回退。这时候不要慌多跑一段时间看整体趋势。如果回退非常严重可以在新环境里先增大探索噪声比如提高 action noise、增大 entropy 系数让策略适应新目标后再恢复原始参数。4.4 reset_num_timesteps 的时间步语义SB3 的 learn() 有一个参数叫reset_num_timesteps默认是 True。意思是每次调用 learn() 时训练步数计数器从 0 重新开始。这个参数在断点续训时很容易被忽略后果是 TensorBoard 里的训练曲线会被割裂成好几段步数都从 0 开始肉眼很难判断整体训练进度是否正常。如果你希望日志里的时间轴保持连续第二次调用 learn() 时设成 Falsemodel.learn(total_timesteps100000, reset_num_timestepsTrue, tb_log_namephase1) model.learn(total_timesteps100000, reset_num_timestepsFalse, tb_log_namephase2)这样第二个阶段的日志步数会从 100000 开始累计训练曲线连续可读。我在实战中会用tb_log_name区分不同阶段并配合reset_num_timesteps控制语义这比后期从 TensorBoard 里手动修补要省事得多。5. 完整案例CartPole 从训练、存档到二次提升5.1 环境准备与依赖安装直接给一套我常用的依赖组合SB3 2.x gymnasium 是当前最稳定的搭配pip install stable-baselines32.2.1 gymnasium torch如果你只是做入门实验CPU 跑 CartPole 完全够用。安装完成后先验证一下环境输出是不是 5 元组import gymnasium as gym env gym.make(CartPole-v1) obs, info env.reset() print(obs, info) print(env.action_space, env.observation_space)输出Box([...])和Discrete(2)就对了。如果 env.reset() 只返回 obs说明你装的是旧版 gym需要升级或换用 gymnasium。5.2 第一阶段训练与保存用下面这段代码训练一个 10 万步的 PPO 模型并同时保存模型和检查点from stable_baselines3 import PPO from stable_baselines3.common.callbacks import CheckpointCallback import gymnasium as gym env gym.make(CartPole-v1) model PPO( MlpPolicy, env, verbose1, learning_rate3e-4, n_steps2048, batch_size64, gamma0.99, gae_lambda0.95, clip_range0.2, ) checkpoint_callback CheckpointCallback( save_freq20000, save_path./checkpoints/, name_prefixppo_cartpole, ) model.learn(total_timesteps100000, callbackcheckpoint_callback) model.save(ppo_cartpole_final)训练完成后./checkpoints/下应该能看到 5 个存档点。这个阶段的目标不是跑出多漂亮的奖励而是验证整个存档链路是否通畅。5.3 第二阶段读取模型做推理评估把模型加载回来用 deterministic 策略跑 100 个 episode统计一下平均回报import gymnasium as gym import numpy as np from stable_baselines3 import PPO env gym.make(CartPole-v1) model PPO.load(ppo_cartpole_final, envenv) episode_rewards [] for _ in range(100): obs, info env.reset() ep_rew 0.0 terminated truncated False while not (terminated or truncated): action, _ model.predict(obs, deterministicTrue) obs, reward, terminated, truncated, info env.step(action) ep_rew reward episode_rewards.append(ep_rew) print(np.mean(episode_rewards), np.std(episode_rewards))如果平均奖励接近 500说明策略已经能稳定撑满整个 episodeCartPole 这个入门任务就宣告通关了。在实际项目中这一步我会包成一个独立的evaluate.py每次训练后都跑一遍作为模型是否达到预期指标的验收闸门。5.4 第三阶段加载再训练把奖励上限再顶上去CartPole-v1 的最大 episode 长度是 500第一阶段模型很可能已经接近上限。这种已经收敛的任务里再训练通常不是为了提升成绩而是测试再训练链路是否稳。可以故意用低一个档次的初始模型做实验先训练 2 万步保存加载后再训 2 万步对比两次的评估分数确认分数确实有增长而不是原地踏步。env gym.make(CartPole-v1) model PPO.load(checkpoints/ppo_cartpole_20000_steps, envenv) model.learning_rate 1e-4 model.learn(total_timesteps20000, reset_num_timestepsFalse, tb_log_nameretrain) model.save(ppo_cartpole_retrained)对比 TensorBoard 里两个阶段的曲线时只要第二阶段曲线的起点不比第一阶段结束点低太多或者很快就拉回原来水平就说明再训练链路是健康的。如果第二阶段明显下滑且长时间无法恢复那通常不是代码问题而是学习率没调好或者环境配置在加载后和原训练不一致。5.5 扩展思考这套打法怎么用到真实场景很多热门的强化学习方向都可以直接套用这套保存-读取-再训练流程。比如机械臂强化学习实战里通常先花大量时间在仿真器里跑出基础策略保存模型再迁移到真实机械臂上做微调仿真阶段遇到服务器重启靠 CheckpointCallback 保住进度。又比如 MILP 与强化学习的交叉方向用强化学习求解组合优化问题时奖励函数往往要先定义粗略版本跑通流程后续再逐步细化这种迭代天然依赖模型再训练能力。甚至离线强化学习场景比如 IQL 从固定数据集训练出初始策略后再接入在线环境做一小段微调本质上也是读取离线阶段保存的模型 在线小步长再训练。掌握 SB3 这套生命周期管理等于给这些高级玩法打下了地基。6. 实战中常见的坑和排查方法6.1 gym/gymnasium 版本错位这是加载模型报错概率最高的问题。SB3 2.x 已经全面使用 gymnasium如果你环境里同时装了旧的 gym或者训练脚本和加载脚本用了不同的环境库空间检查会直接失败。排查方法很简单在训练和加载脚本里都打印env.reset()的返回结构确认都是(obs, info)的二元组。如果发现一个返回二元组、一个返回旧版三元组优先统一环境库版本。6.2 加载后性能下降的排查思路加载模型用于推理后发现表现远不如训练时先按这个顺序排查第一确认predict是否用了deterministicTrue第二确认是否用了 VecNormalize 且统计量是否恢复第三确认环境本身是否可复现——比如随机种子不同导致初始状态不同第四检查模型文件是不是覆盖了有时候训练后期覆盖了中间更好表现的存档。第四点特别容易踩所以我用 CheckpointCallback 时会把 save_freq 设置得小一些后期再用评估回调挑最优存档。6.3 跨机器恢复、换卡换CPU的注意事项模型从 GPU 训练机上保存拿到没有 GPU 的机器上加载推理通常不需要额外设置PyTorch 会自动做设备映射。但如果你显式指定了devicecuda在没有 GPU 的机器上就会报错。更稳妥的做法是在 load 时不传 device让 SB3 自动判断或者统一写成deviceauto。跨机器还有一个常被忽略的坑自定义策略类的模块路径。A 机器上你的自定义特征提取器在utils.networks.MyFeatureExtractorB 机器上找不到这个模块加载就会失败。跨机器迁移前把自定义网络类打包成模块或写成可安装的包能省去很多麻烦。6.4 从离线强化学习到在线微调IQL 和 SB3 的衔接思路最近经常被问到 IQL 这类离线强化学习训练的模型怎么拿到 SB3 里继续在线训练。严格来说 IQL 是独立实现的不直接用 SB3 训练但思路可以打通离线阶段学出来的 Q 网络或策略权重可以通过自定义 policy 的初始化方式注入到 SB3 的 SAC 或 PPO 里然后在在线环境中小步长再训练。对绝大多数工程场景更实际的做法是先用 SB3 的 SAC 从离线收集到的 replay buffer 数据做若干轮 offline 更新再把同一套模型接入真实环境进行在线微调。核心准备工作就是SAC.load()之后立刻load_replay_buffer()保证离线数据不丢失然后大幅调低学习率防止在线阶段刚一开始就把离线阶段学到的先验知识冲掉。6.5 一个好用的小技巧训练阶段给每个存档点写说明项目久了checkpoints 目录里会躺着几十上百个模型光靠文件名根本分不清哪个对应哪个实验。我习惯在每次 save 之后写一个 json 摘要记录训练步数、学习率、奖励函数版本、环境版本、备注信息。下次选模型时先看 json 再决定加载哪一个。这个习惯帮我避免过好多次加载了半天发现是用旧奖励函数训的这种事故。根据我个人的实战体会模型保存、读取、再训练这套流程单独看每一个 API 都不难难的是一开始就把它当成项目的骨架来设计。项目第一天就把 CheckpointCallback、VecNormalize 保存、replay buffer 存档、日志分阶段这些机制搭好后面所有实验都会顺畅很多否则等到训练了三天才发现中间环节缺失那才是最尴尬的时刻。建议你也从今天的小实验开始把会训练升级成会管理训练。