训练一个强化学习模型,最怕的是什么?不是收敛慢,是辛辛苦苦跑了一晚上的实验,第二天发现没保存,全白干了。我刚上手 Stable Baselines3(SB3)那会儿,踩的第一个大坑就是这个——以为 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, verbose=1) model.learn(total_timesteps=50000) 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() 挂上 CheckpointCallback:
from stable_baselines3.common.callbacks import CheckpointCallback checkpoint = CheckpointCallback( save_freq=10000, save_path="./checkpoints/", name_prefix="ppo_cartpole", save_vecnormalize=True, ) model.learn(total_timesteps=200000, callback=checkpoint)这样每 10000 步就会自动生成一个带时间步标记的存档,比如ppo_cartpole_10000_steps.zip、ppo_cartpole_20000_steps.zip。训练中途崩溃,损失最多只占一个 save_freq 的间隔;训练结束后想回退到某个中间状态,也可以直接挑一个时间拾取点加载,非常实用。
save_vecnormalize=True这个参数值得单独说一下。如果训练脚本里用了 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_obs=True, norm_reward=True, clip_obs=10.0) model = PPO("MlpPolicy", venv, verbose=1) model.learn(total_timesteps=100000) 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", env=venv)如果你是用来继续训练,记得把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_format="pth")但 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, verbose=1) model.learn(total_timesteps=50000) model.save("sac_pendulum") model.save_replay_buffer("sac_replay_buffer") # 恢复时 model = SAC.load("sac_pendulum", env=env) model.load_replay_buffer("sac_replay_buffer") model.learn(total_timesteps=50000)注意:
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, deterministic=True) obs, reward, terminated, truncated, info = env.step(action) if terminated or truncated: obs, info = env.reset()这里有两个细节。第一,predict(obs, deterministic=True)表示选择确定性动作,即策略网络输出的均值或 argmax;如果设成 False,则会从动作分布中采样,适合训练或探索阶段。二是 predict 返回的是一个二元组(action, state),第二个元素是 RNN/GRU 等循环策略使用的隐藏状态,普通 MLP 策略直接忽略即可。
如果你要接着训练,需要在 load 时把环境传进去:
env = gym.make("CartPole-v1") model = PPO.load("ppo_cartpole", env=env) model.learn(total_timesteps=30000)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_dim=256): 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_class=MyFeatureExtractor) model = PPO("MlpPolicy", env, policy_kwargs=policy_kwargs) model.learn(20000) model.save("ppo_custom_extractor") # 在另一个脚本中加载 model = PPO.load( "ppo_custom_extractor", env=env, 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", env=env) # 第二步:继续训练,这里的 total_timesteps 是增量,不是累计目标 model.learn(total_timesteps=100000) # 第三步:重新保存 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", env=env) model.learning_rate = 1e-4 model.learn(total_timesteps=50000)如果你的学习率是通过 schedule 函数传入的,比如learning_rate=lambda 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() 时设成 False:
model.learn(total_timesteps=100000, reset_num_timesteps=True, tb_log_name="phase1") model.learn(total_timesteps=100000, reset_num_timesteps=False, tb_log_name="phase2")这样第二个阶段的日志步数会从 100000 开始累计,训练曲线连续可读。我在实战中会用tb_log_name区分不同阶段,并配合reset_num_timesteps控制语义,这比后期从 TensorBoard 里手动修补要省事得多。
5. 完整案例:CartPole 从训练、存档到二次提升
5.1 环境准备与依赖安装
直接给一套我常用的依赖组合,SB3 2.x + gymnasium 是当前最稳定的搭配:
pip install stable-baselines3==2.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, verbose=1, learning_rate=3e-4, n_steps=2048, batch_size=64, gamma=0.99, gae_lambda=0.95, clip_range=0.2, ) checkpoint_callback = CheckpointCallback( save_freq=20000, save_path="./checkpoints/", name_prefix="ppo_cartpole", ) model.learn(total_timesteps=100000, callback=checkpoint_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", env=env) 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, deterministic=True) 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,说明策略已经能稳定撑满整个 episode,CartPole 这个入门任务就宣告通关了。在实际项目中,这一步我会包成一个独立的evaluate.py,每次训练后都跑一遍,作为模型是否达到预期指标的验收闸门。
5.4 第三阶段:加载再训练,把奖励上限再顶上去
CartPole-v1 的最大 episode 长度是 500,第一阶段模型很可能已经接近上限。这种已经收敛的任务里,再训练通常不是为了提升成绩,而是测试"再训练链路"是否稳。可以故意用低一个档次的初始模型做实验:先训练 2 万步,保存;加载后再训 2 万步,对比两次的评估分数,确认分数确实有增长而不是原地踏步。
env = gym.make("CartPole-v1") model = PPO.load("checkpoints/ppo_cartpole_20000_steps", env=env) model.learning_rate = 1e-4 model.learn(total_timesteps=20000, reset_num_timesteps=False, tb_log_name="retrain") 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是否用了deterministic=True;第二,确认是否用了 VecNormalize 且统计量是否恢复;第三,确认环境本身是否可复现——比如随机种子不同导致初始状态不同;第四,检查模型文件是不是覆盖了,有时候训练后期覆盖了中间更好表现的存档。第四点特别容易踩,所以我用 CheckpointCallback 时会把 save_freq 设置得小一些,后期再用评估回调挑最优存档。
6.3 跨机器恢复、换卡换CPU的注意事项
模型从 GPU 训练机上保存,拿到没有 GPU 的机器上加载推理,通常不需要额外设置,PyTorch 会自动做设备映射。但如果你显式指定了device="cuda",在没有 GPU 的机器上就会报错。更稳妥的做法是在 load 时不传 device,让 SB3 自动判断;或者统一写成device="auto"。跨机器还有一个常被忽略的坑:自定义策略类的模块路径。A 机器上你的自定义特征提取器在utils.networks.MyFeatureExtractor,B 机器上找不到这个模块,加载就会失败。跨机器迁移前,把自定义网络类打包成模块或写成可安装的包,能省去很多麻烦。
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 存档、日志分阶段这些机制搭好,后面所有实验都会顺畅很多;否则等到训练了三天才发现中间环节缺失,那才是最尴尬的时刻。建议你也从今天的小实验开始,把"会训练"升级成"会管理训练"。