十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

SB3模型保存与再训练实战指南:从checkpoint到续训

SB3模型保存与再训练实战指南:从checkpoint到续训 训练到一半断电大概是每个搞强化学习的人都经历过的噩梦。我用Stable Baselines3跑一个机械臂仿真任务连续跑了快两天结果一次意外重启让所有训练进度清零。当时脑子里只有一个念头为什么没早点把模型保存、读取、再训练这条链路彻底搞清楚。后来我把SB3里这套机制完整摸了一遍发现官方文档给的示例太简略很多坑藏在实际操作里——模型加载后策略乱动、继续训练loss直接崩掉、自定义环境报错找不到类。这篇文章就把我整理的完整流程和经验写出来希望能帮你少走弯路。不管你是刚接触强化学习入门的新手还是已经在跑PPO、DQN、IQL这类算法的老手只要涉及长时间训练、模型评估或迁移实验都逃不开模型的保存和再训练。内容不限定单一算法以Stable Baselines3的PPO为例原理同样适用于其他内置算法。1. 为什么我会把保存/读取练成肌肉记忆先说观点在Stable Baselines3里model.save()和model.load()不是训练完之后的收尾动作而是实验一开始就要设计进去的机制。很多人只在最后保存一次模型这种做法一旦中间出问题前面的算力全部白费。1.1 保存模型不只是防中断更是实验管理的基础我之前做过一组超参数对比实验分别用不同学习率训练同一个机械臂环境。如果每个实验只在结束时保存那么我想对比训练中期和后期策略差异时完全没有办法拿到中间状态。后来我在训练脚本里加了一个每隔固定步数自动保存的机制对比分析时才有了完整的数据支撑。除了防中断和对比实验保存模型还有三个常见用途评估中间策略判断当前策略是否还有继续训练的价值如果评估曲线已经平台期就可以提前停止避免浪费时间。共享初始化在做域随机化或环境微调时用同一个预训练模型作为起点能大幅缩短新任务的收敛时间。部署与回滚训练出的最佳模型需要存档如果后续实验把模型搞坏了还能从之前保存的版本恢复。1.2 保存频率怎么选这里有个容易忽略的细节SB3官方提供的CheckpointCallback可以很方便地周期性保存模型但它有一个容易踩的坑save_freq的单位是timesteps不是episodes。如果你设置的save_freq小于算法单次更新的步数n_steps那么保存动作会发生在一次更新还没结束的中间状态虽然不会报错但保存下来的模型对应的优化器状态和真实训练状态可能不完全同步。我自己用的配置是训练总步数n_stepssave_freq说明10万左右20485000中等频率适合调试100万以上204850000低频保存减少IO消耗长期任务4096100000配合早期停止策略如果你的训练任务本身很短比如几万步以内我建议直接在learn()结束保存一份再加一个手动评估节点就够了没必要搞checkpoint目录。训练时间越长checkpoint策略越重要。2. save/load这组API的真实工作方式SB3的save()和load()从接口上看简单得吓人一行就能完成但很多人不知道这背后到底保存了什么、加载时要满足什么条件。2.1 model.save()到底保存了哪些内容调用model.save(ppo_cartpole)后SB3会生成一个zip压缩包如果没写扩展名会自动加.zip。里面不只是神经网络的权重还包括策略网络的PyTorch权重policy.pth优化器状态policy.optimizer.pth算法超参数学习率、gamma、gae_lambda等环境ID和sample的observation space / action space信息系统信息文本这意味着你在加载模型时不需要手动重新声明policy_kwargs网络结构会从保存的配置里恢复。但这也带来一个限制模型文件强依赖原来的环境和策略类定义。如果换了环境类名或算法版本反序列化会失败。2.2 load()的常见参数与使用误区PPO.load()最基础用法是import gymnasium as gym from stable_baselines3 import PPO env gym.make(CartPole-v1) model PPO.load(ppo_cartpole.zip, envenv)注意加载时传入的env不是必须的。如果你只做纯推理不调用learn()不传env也能加载。但如果你要继续训练或者要调用evaluate_policy做评估就必须传入一个和训练时等价的env对象。这里有个高频错误使用A2C.load()去加载一个PPO训练出来的zip包会直接报错或产生不可预料的网络结构问题。因为保存时记录的是policy_class不同算法对应不同策略网络跨算法加载没有意义。一定要保证加载类与保存类一致。load()还有一个容易忽略的参数custom_objects。当保存的模型引用了自定义类比如你写了一个自定义的feature_extractor或回调对象加载时如果找不到类定义会报AttributeError。此时可以用custom_objects参数手动指定替代对象。不过更稳妥的做法是把自定义类写在独立模块里确保训练和加载时都能import到。2.3 保存目录管理与命名规范我见过不少同学把所有模型都命名为model.zip训练两轮之后根本分不清哪个是哪个。这里分享一个我长期在用的命名规则import os import time from stable_baselines3 import PPO from stable_baselines3.common.callbacks import CheckpointCallback save_dir f./checkpoints/{time.strftime(%Y%m%d_%H%M%S)} os.makedirs(save_dir, exist_okTrue) checkpoint_callback CheckpointCallback( save_freq50000, save_pathsave_dir, name_prefixppo_mechanical_arm ) model PPO(MlpPolicy, env, verbose1) model.learn(total_timesteps1_000_000, callbackcheckpoint_callback)CheckpointCallback会自动生成类似ppo_mechanical_arm_50000_steps.zip这样的文件加上顶层目录的时间戳训练历史一目了然。我还会额外记录一份metadata.json把环境ID、总步数、学习率、网络结构等信息都写进去方便日后复盘。3. 再训练最容易翻车的地方环境、学习率与归一化模型加载成功不等于你会再训练。model.learn()续跑时有三个地方最容易翻车而且翻车现场往往非常隐蔽。3.1 环境一致性是续训的命门SB3的模型在保存时记录了原始环境的ID比如CartPole-v1load()之后你可以通过model.get_env()重新拿到环境对象。但这里有个陷阱如果你训练时用了VecNormalize向量环境归一化只保存模型不保存VecNormalize的状态加载后再训练时观测和奖励的均值和方差是乱的策略相当于换了一套输入分布表现会断崖式下跌。正确做法是把env状态单独保存from stable_baselines3.common.vec_env import VecNormalize from stable_baselines3.common.vec_env import make_vec_env env make_vec_env(CartPole-v1, n_envs4) env VecNormalize(env, norm_obsTrue, norm_rewardTrue) model PPO(MlpPolicy, env, n_steps2048) model.learn(total_timesteps200_000) env.save(vec_normalize.pkl) model.save(ppo_cartpole_vec)再训练时必须用VecNormalize.load()把归一化参数恢复env make_vec_env(CartPole-v1, n_envs4) env VecNormalize.load(vec_normalize.pkl, env) model PPO.load(ppo_cartpole_vec.zip, envenv)注意顺序不能反先加载vec_normalize.pkl再把它作为env传给PPO.load()。如果你把两个文件放在同一目录建议统一命名避免后期不知道哪个pkl对应哪个模型。3.2 学习率调度器在续训时可能毁掉已有策略这是再训练时最隐蔽的坑。SB3的PPO在初始化时通过learning_rate参数构建了一个学习率调度器保存模型时当前调度器的状态也会被保存。当你在load()之后直接调用learn()继续训练时调度器会按照保存时的参数重新开始调度。如果原来的learning_rate给得比较大比如0.0003继续训练时一开始就用这么大的步长去更新已有的策略极容易造成策略剧烈波动表现为评估分数先猛涨再猛跌或者直接发散。我的经验是续训练习率一般设为原学习率的1/3到1/10。比如原训练用learning_rate3e-4续训可以改成1e-4或3e-5。SB3的load()允许直接覆盖原来的超参数model PPO.load( ppo_cartpole_vec.zip, envenv, learning_rate3e-5 )这样会覆盖原来保存的学习率效果等于用一组更保守的更新步长微调已有策略。对于已经训练得比较充分的模型我一般用一个很小的固定学习率比如constant_fn(1e-5)避免中期再次出现剧烈波动。3.3 reset_num_timesteps到底该怎么设learn(total_timesteps100_000, reset_num_timestepsFalse)里reset_num_timestepsFalse表示不重置当前的总步数计数。这个参数直接影响学习率调度器的位置。如果训练脚本里把reset_num_timesteps设成True调度器会回到初始状态可能把学习率重新拉高再次踩中上面说的坑。我通常这样处理训练阶段第一次调用learn()时用默认的True后续所有续跑都显式设置为False。这样训练日志里的time/total_timesteps能连续累积学习率调度也维持原状只有在你故意覆盖学习率时才不受影响。4. 三个让我debug到怀疑人生的坑含完整排查链路下面三个问题我都实际遇到过每一个都花了不少时间定位。这里我把排查思路完整写出来而不是直接给结论因为排查方法本身比答案值钱。4.1 坑一加载后的模型表现像个随机策略第一次出现这个现象时我用evaluate_policy评估保存的模型score从几百掉到个位数第一反应是“保存坏了”。我重新训练了一版并立刻评估发现保存前score是好的加载后score奇差于是开始怀疑load()写错了。排查链路是这样的第一步确认load()返回的模型类型。打印type(model)确认是PPO而不是其他类。第二步确认评估时用的env和训练时一致。结果发现手动测试时我创建了一个CartPole-v1普通环境但训练时用的是VecNormalize包装过的环境。问题就出在这——训练时观测被归一化手动评测时观测没被归一化模型看到的输入分布完全不同。修复方式是加载vec_normalize.pkl并用它包住测试环境env make_vec_env(CartPole-v1) env VecNormalize.load(vec_normalize.pkl, env) env.training False # 评估时不需要更新running mean/var env.norm_reward False # 评估时一般只关心原始reward这里有必要提一下env.training False它告诉VecNormalize在评估阶段不要再更新内部的均值和方差。否则评估一遍统计量就被改变一次结果不稳定。4.2 坑二继续训练时loss直接NaN这个坑最让人头大。一个训练正常的模型续跑几百步后loss变成NaN然后整个策略就废了。最初我以为是学习率太大调低之后依然NaN于是开始排查环境。因为我的任务里用到了Box连续动作空间怀疑action有异常值。我在环境step里加了断言发现动作范围没问题。后来我把问题定位到优化器状态。SB3在保存模型时确实保存了优化器状态但如果你保存的模型是训练中期保存的优化器中可能存在历史梯度状态。续跑时如果再叠加较大的学习率极少数情况下可能导致梯度爆炸。换成learning_rate1e-5之后NaN问题消失但这样治标不治本。最终我查看了优化器参数里是否有NaN发现触发点其实是自定义奖励函数里出现了一次inf——某个中间量除零了。策略本身没问题是奖励信号瞬间异常导致loss冲高。修复奖励函数后重新训练问题彻底消失。这个坑的排查思路值得记住遇到NaN先看环境奖励输出再看网络输出最后才怀疑优化器和学习率。顺序反了会浪费大量时间。4.3 坑三自定义环境在加载时报错找不到类SB3保存模型时会记录环境ID如果你用的是自定义环境MyEnv-v0加载时没有注册就会报类似MLP policy... error: Environment with ID MyEnv-v0 not found。我之前写过一版环境用了gym.register注册但注册代码写在训练脚本里。load()时只import了算法库没有执行注册函数于是报错。解决办法有两种把环境定义写在独立Python模块中加载前先import my_env触发注册。直接创建好env传给load()这样SB3不会尝试通过环境ID自动重建环境。第二种方法更省心from my_env_module import make_env env make_env(MyEnv-v0) model PPO.load(my_model.zip, envenv)如果你用了make_vec_env也是一样的逻辑重点是确保训练、评估、再训练三处拿到的环境实例在观测空间、动作空间和内部逻辑上完全等价。5. 一套可以抄走的checkpoint工作流与续训技巧前面把原理和坑讲完了下面给出一套我自己长期在用的工作流。你直接按这个结构改就能跑起来。5.1 训练脚本骨架自动保存定时评估import gymnasium as gym import numpy as np from stable_baselines3 import PPO from stable_baselines3.common.callbacks import ( CheckpointCallback, EvalCallback, CallbackList ) from stable_baselines3.common.env_util import make_vec_env from stable_baselines3.common.vec_env import VecNormalize # 创建并行环境 归一化 env make_vec_env(CartPole-v1, n_envs4) env VecNormalize(env, norm_obsTrue, norm_rewardTrue) # 自动保存 checkpoint_callback CheckpointCallback( save_freq20000, save_path./checkpoints/run1, name_prefixppo_cartpole ) # 定时评估保存best model eval_env make_vec_env(CartPole-v1, n_envs1) eval_env VecNormalize(eval_env, norm_obsTrue, norm_rewardTrue) eval_callback EvalCallback( eval_env, best_model_save_path./checkpoints/run1/best_model, log_path./checkpoints/run1/eval_logs, eval_freq5000, n_eval_episodes10, deterministicTrue ) model PPO( MlpPolicy, env, n_steps2048, batch_size64, n_epochs10, learning_rate3e-4, verbose1 ) model.learn( total_timesteps200_000, callbackCallbackList([checkpoint_callback, eval_callback]) ) env.save(./checkpoints/run1/vec_normalize.pkl) model.save(./checkpoints/run1/final_model.zip)EvalCallback的best_model_save_path会自动保存评估分数最高的模型命名是best_model.zip。如果你经常做模型筛选这个机制比手动记录epoch数省心得多。5.2 续训脚本骨架低学习率不重置步数import gymnasium as gym from stable_baselines3 import PPO from stable_baselines3.common.env_util import make_vec_env from stable_baselines3.common.vec_env import VecNormalize from stable_baselines3.common.callbacks import CheckpointCallback # 重新创建同样配置的环境 env make_vec_env(CartPole-v1, n_envs4) env VecNormalize.load(./checkpoints/run1/vec_normalize.pkl, env) model PPO.load( ./checkpoints/run1/best_model/best_model.zip, envenv, learning_rate3e-5 # 覆盖为一较小的固定值 ) checkpoint_callback CheckpointCallback( save_freq20000, save_path./checkpoints/run2, name_prefixppo_cartpole ) model.learn( total_timesteps100_000, reset_num_timestepsFalse, callbackcheckpoint_callback ) env.save(./checkpoints/run2/vec_normalize.pkl) model.save(./checkpoints/run2/continued_model.zip)续训时最核心的就是三件事恢复vec_normalize.pkl、降低学习率、reset_num_timestepsFalse。5.3 一个小技巧用metadata隔离不同训练阶段如果你和我一样常做多阶段训练建议在每个checkpoint目录里放一份metadata.json至少包含算法名、环境ID、总步数、学习率、评估分数、备注。import json import time metadata { algorithm: PPO, env_id: CartPole-v1, total_timesteps: 200000, learning_rate: 3e-4, note: 第一次训练baseline, saved_at: time.strftime(%Y-%m-%d %H:%M:%S) } with open(./checkpoints/run1/metadata.json, w) as f: json.dump(metadata, f, indent2)别小看这个动作。等你过了一个月再回来看一堆ppo_cartpole_50000_steps.zip没有metadata基本等于没保存。有了它你还能写个小脚本自动读取每个run的评估分数直接列成表格做对比。5.4 什么时候该续训什么时候该重新开始最后聊一个经常被问的问题评估曲线进入平台期后是续训还是重新跑我的判断依据很简单如果最近几次评估分数的均值还在缓慢上升说明策略还有空间续训有效如果已经来回震荡且连续多个checkpoint没有刷新最高分再续训大概率只是浪费时间此时应该去改环境奖励、网络结构或探索参数而不是单纯堆步数。我自己在机械臂任务里试过同一份预训练模型续训20万步前15万步分数稳定提升后5万步反而下滑。后来发现是学习率没调低策略在最优点附近来回冲撞。把学习率降到原来的1/5后平台期才真正稳住。这个经验适用于大多数连续动作任务续训的步长必须比首次训练更小别让调度器把策略从已经学好的区域推出去。保存、读取、再训练这套流程说白了就是给强化学习实验做版本管理。你不需要一开始就写得多完美但至少要把checkpoint和VecNormalize这两个事刻在脑子里它们能帮你避开90%的续训翻车现场。
返回列表