强化学习智能体怎么保存?用 Stable-Baselines3 保存并核对载入模型

保存 PPO 模型与环境元信息,重新载入后比较同一观察的动作。提供演示数据、核对方法和适用边界。

保存强化学习智能体,至少要保留模型文件和重建输入输出所需的环境信息。文件写入成功只能说明保存动作完成,重新载入并核对同一观察的动作,才是检查模型能否继续使用的下一步。下面用 Stable-Baselines3 的 PPO 演示这个过程。

准备隔离环境与保存目录

在独立 Python 环境安装 stable-baselines3 与 gymnasium,并记录实际版本。例子使用无画面窗口的 CartPole-v1,不会控制真实设备。代码只训练一小段时间来演示保存流程,不承诺学会平衡小车。

强化学习智能体怎么保存?用 Stable-Baselines3 保存并核对载入模型

把下面代码保存为 save_ppo.py。如果目录里已有重要模型,换一个新的演示目录再运行,避免用同名文件覆盖它。若已创建并启用了独立 Python 环境,可用以下命令安装依赖;后面保存完程序再运行第二条命令。

python -m pip install stable-baselines3 gymnasium
python save_ppo.py
import json
from pathlib import Path
from importlib.metadata import version
import numpy as np
import gymnasium as gym
from stable_baselines3 import PPO

folder = Path("ppo_save_demo")
folder.mkdir(exist_ok=True)
if (folder / "agent.zip").exists():
    raise FileExistsError("请改用新目录,避免覆盖原模型")

env = gym.make("CartPole-v1")
model = PPO("MlpPolicy", env, seed=17, n_steps=256, batch_size=64, verbose=0)
model.learn(total_timesteps=1024)

observations = []
obs, info = env.reset(seed=18)
for _ in range(20):
    observations.append(obs.copy())
    action, _ = model.predict(obs, deterministic=True)
    obs, reward, terminated, truncated, info = env.step(int(action))
    if terminated or truncated:
        obs, info = env.reset()
observations = np.asarray(observations)
before, _ = model.predict(observations, deterministic=True)
model.save(str(folder / "agent"))
metadata = {
    "algorithm": "PPO", "environment": "CartPole-v1",
    "gymnasium": version("gymnasium"),
    "stable_baselines3": version("stable-baselines3"),
    "torch": version("torch"),
    "observation_space": str(env.observation_space),
    "action_space": str(env.action_space)
}
(folder / "metadata.json").write_text(
    json.dumps(metadata, ensure_ascii=False, indent=2), encoding="utf-8")
loaded = PPO.load(str(folder / "agent"), env=env)
after, _ = loaded.predict(observations, deterministic=True)
assert np.array_equal(before, after), (before, after)
print("同一组观察的离散动作一致:", len(observations))
env.close()

逐项核对保存结果

运行后先检查 agent.zip 与 metadata.json 是否存在,再检查断言是否通过。这里比较的是相同观察上的确定性预测,避免把两个不同回合的随机环境状态混到一起。离散动作一致是一个有限样本检查,不是完整网络等价证明,也不是训练效果评测。

载入调用要接收返回值,例如 loaded = PPO.load(...)。不要先新建对象再调用其 load(),却继续使用原对象;这容易误以为参数已经被替换。

模型文件以外还应保留什么

如果实际项目使用了观察归一化、奖励处理、自定义环境或动作映射,单独保存这些配置及相应统计量,并在载入时按原顺序重建。上面的例子没有这些附加处理。还应记录训练代码版本、随机种子和评估场景,使你知道模型原来接受什么输入。

准备续训时,另外核对优化器状态、算法参数和总时间步数的处理;如果换成带经验回放的算法,还应按该算法文档确认是否另存回放缓冲区。不能从“有一个 zip”推出所有训练过程都能精确恢复。

载入失败怎么排查

  1. 先检查文件路径和所用算法是否一致,PPO 保存的模型用 PPO 载入。
  2. 检查环境的观察与动作空间是否改变;不要为了绕过错误硬改维度。
  3. 比较保存与载入环境的包版本,可使用 PPO.load(path, print_system_info=True) 辅助查看系统差异。
  4. 只载入可信来源的模型文件。此类序列化格式可能包含复杂对象,不能把陌生文件当成普通数据随意反序列化。

本篇未运行 PPO 训练与保存载入;仅核对代码语法和官方接口,完整流程需在独立环境执行后验收。示例训练步数没有经过任务效果调参。跨版本恢复和自定义环境兼容性需另行验证。

接口与保存边界依据 2026 年 10 月 3 日读取的 SB3 保存格式说明、SB3 保存载入示例及 Gymnasium 基本交互接口。文档 master 分支会更新,使用时应切换到已安装版本对应的说明。

Ai菜鸟网。发布者:AI小管家,转载请注明出处:https://www.alyyhw.com/32958.html

赞 (0)
AI小管家的头像AI小管家
强化学习中的状态是什么?用移动小车说明状态与观察的区别
上一篇 1天前
多智能体强化学习的 Q 值代表什么?分清自身动作与联合动作
下一篇 1天前

相关推荐

联系我们

联系我们

1

在线咨询: QQ交谈

邮件:admin@example.com

工作时间:周一至周五,9:30-18:30,节假日休息

关注微信
关注微信
分享本页
返回顶部