Stable-Baselines3 快速上手:用 A2C 训练并运行你的第一个强化学习智能体
【免费下载链接】stable-baselines3PyTorch version of Stable Baselines, reliable implementations of reinforcement learning algorithms.项目地址: https://gitcode.com/GitHub_Trending/st/stable-baselines3
Stable-Baselines3(SB3)是一套基于 PyTorch 的可靠强化学习算法实现库,其所有算法共享一套 sklearn 风格(fit/predict式)的统一接口。本篇基于仓库官方快速上手文档 docs/guide/quickstart.md 展开,结合源码带你从零完成第一个训练任务:在 CartPole-v1 上用 A2C 训练一个智能体、用训练好的模型做推理与可视化,并理解训练背后由向量化环境(VecEnv)驱动的数据流。
SB3 的核心设计:一套接口,全部算法
SB3 中所有强化学习算法(A2C、PPO、DQN、SAC、TD3、DDPG)都遵循统一的 sklearn 风格语法:构造模型 →learn()训练 →predict()推理。这一设计在 stable_baselines3/init.py 中通过统一的顶层导出体现,六个算法类共用相同的 API 形态。
在动手写代码前,有一个必须理解的前提:SB3 内部使用向量化环境(VecEnv)而不是单个 Gym 环境。也就是说,即使你只传入一个环境,框架也会把它包装成"同时运行 1 个环境副本"的向量环境。关于 VecEnv 的完整特性与它和单个 Gym 环境的差异,请阅读仓库文档 docs/guide/vec_envs.md,这里先记住三个最关键的差异:
vec_env.reset()只返回观测(obs),不返回 Gym 0.26+ 的(obs, info)元组,reset 时的信息存于vec_env.reset_infos;vec_env.step(actions)接收批量动作,返回四元组obs, rewards, dones, infos(而非 Gym 的五元组),其中dones = terminated or truncated,且obs/rewards/dones都是带 batch 维度的 NumPy 数组;- VecEnv 会在每个回合结束时自动重置环境,因此
done[i]为 True 时返回的观测其实是下一回合的首个观测;若需要被终止回合的"真实最后观测",请从infos[i]["terminal_observation"]中读取。
从源码看,这种包装发生在 stable_baselines3/common/base_class.py 的_wrap_env()方法中:当你传入普通 Gym 环境时,SB3 会依次为其包上Monitor(记录回合回报/长度)、DummyVecEnv(向量化),对图像观测还会自动包上VecTransposeImage以调整通道顺序。
第一个训练示例:在 CartPole-v1 上训练 A2C
下面的代码来自 docs/guide/quickstart.md,它完整演示了"构造 → 训练 → 用 VecEnv 推理"的全流程:
import gymnasium as gym from stable_baselines3 import A2C env = gym.make("CartPole-v1", render_mode="rgb_array") model = A2C("MlpPolicy", env, verbose=1) model.learn(total_timesteps=10_000) vec_env = model.get_env() obs = vec_env.reset() for i in range(1000): action, _state = model.predict(obs, deterministic=True) obs, reward, done, info = vec_env.step(action) vec_env.render("human") # VecEnv resets automatically # if done: # obs = vec_env.reset()逐步拆解这段代码:
gym.make("CartPole-v1", render_mode="rgb_array"):创建 Gymnasium 环境,并指定rgb_array渲染模式。SB3 官方推荐使用rgb_array,因为它既能以 NumPy 数组形式取回图像,又能通过 OpenCV 以"human"模式弹窗显示(见 docs/guide/vec_envs.md 中关于渲染的说明)。A2C("MlpPolicy", env, verbose=1):以策略名"MlpPolicy"构造 A2C 模型。verbose=1会在训练时打印设备信息、包装器使用情况和训练日志。model.learn(total_timesteps=10_000):训练 10000 个时间步。model.get_env():取回模型内部的(向量化)环境,用于后续手动推理。- 循环推理:
model.predict(obs, deterministic=True)使用确定性动作(即取均值/最大概率动作而非采样);vec_env.step(action)推进环境并返回四元组。注释特别提醒:VecEnv 自动重置,无需手动调用reset()。
A2C 构造函数核心参数
A2C的定义位于 stable_baselines3/a2c/a2c.py,其构造函数参数及默认值如下:
| 参数 | 默认值 | 说明 |
|---|---|---|
policy | 必填 | 策略别名:"MlpPolicy"/"CnnPolicy"/"MultiInputPolicy"(或策略类本身) |
env | 必填 | 可传入环境实例,或 Gymnasium 中已注册的环境名字符串 |
learning_rate | 7e-4 | 学习率,也可以是"当前训练进度剩余量"(1→0)的调度函数 |
n_steps | 5 | 每次更新前每个环境采样的步数,即单次更新批大小为n_steps * n_env |
gamma | 0.99 | 折扣因子 |
gae_lambda | 1.0 | GAE(广义优势估计)的偏差-方差权衡系数,取 1 时退化为经典优势估计 |
ent_coef | 0.0 | 损失中的熵系数 |
vf_coef | 0.5 | 价值函数损失系数 |
max_grad_norm | 0.5 | 梯度裁剪上限 |
rms_prop_eps | 1e-5 | RMSProp 的 epsilon |
use_rms_prop | True | 是否使用 RMSProp(原始实现)而非 Adam 作为优化器 |
use_sde | False | 是否使用广义状态依赖探索(gSDE)替代动作噪声探索 |
sde_sample_freq | -1 | 使用 gSDE 时每隔多少步重采样噪声矩阵,-1表示仅在 rollout 开始时采样 |
normalize_advantage | False | 是否对优势估计做归一化 |
stats_window_size | 100 | 日志统计窗口:用于平均最近的回合回报、长度等 |
tensorboard_log | None | TensorBoard 日志目录(None表示不记录) |
policy_kwargs | None | 传给策略的额外参数(网络结构、激活函数等) |
verbose | 0 | 0 无输出,1 打印设备/包装器等信息,2 打印调试信息 |
seed | None | 随机种子 |
device | "auto" | 计算设备,auto表示有 GPU 则用 GPU |
值得注意的两点源码细节:
- 策略别名机制:
policy_aliases类属性(见 stable_baselines3/a2c/a2c.py)将"MlpPolicy"映射到ActorCriticPolicy、"CnnPolicy"映射到ActorCriticCnnPolicy、"MultiInputPolicy"映射到MultiInputActorCriticPolicy,并通过BaseAlgorithm._get_policy_from_name()(stable_baselines3/common/base_class.py)解析。因此所有算法都支持同样的三个策略名,切换算法时无需改策略代码。 - 动作空间约束:A2C 仅支持
Box、Discrete、MultiDiscrete、MultiBinary四类动作空间,传入其他类型会直接断言报错。
learn()方法的完整签名
从 stable_baselines3/a2c/a2c.py 中learn()的定义可以看到:
model.learn( total_timesteps=10_000, # 总训练步数(预算) callback=None, # 训练回调(如评估、保存模型) log_interval=100, # 每隔多少步打印一次日志 tb_log_name="A2C", # TensorBoard 运行名 reset_num_timesteps=True, # 连续调用 learn() 时是否重置时间步计数 progress_bar=False, # 是否用 tqdm/rich 显示进度条 )上面的训练循环实际上就是model.learn()内部"收集经验 → 更新策略"的循环封装:collect_rollouts()用当前策略采集轨迹写入 rollout buffer,达到n_steps * n_env步后调用一次train()做一步梯度更新(对 A2C 而言每次更新使用全部数据,见 stable_baselines3/a2c/a2c.py 中train()里for rollout_data in self.rollout_buffer.get(batch_size=None)的单次循环),循环直至总步数达到预算。
用训练好的模型推理与可视化
训练完成后,model.get_env()返回模型训练时使用的同一个 VecEnv(BaseAlgorithm.get_env(),见 stable_baselines3/common/base_class.py),你可以像示例中那样直接驱动它做 rollout 演示。
model.predict()的底层实现在 stable_baselines3/common/policies.py 的BasePolicy.predict()中:
- 传入
deterministic=True时返回确定性动作(策略均值 / 最大概率动作),默认deterministic=False则从分布中采样(保留探索); - 推理在
th.no_grad()下进行,并将动作转回 NumPy; - 对连续动作空间(
Box),超出边界时会被自动裁剪到[low, high]; - 若传入的是单个(非向量化)观测,会自动去掉 batch 维度返回单动作。
一个常见的 API 混用错误值得警惕:如果你把 Gym 的obs, info = env.reset()结果(元组)传给predict(),predict()会抛出明确的ValueError提示"你混淆了 Gym API 与 SB3 VecEnv API"——这也是上面示例中始终坚持用vec_env.reset()(只返回 obs)的原因。
一行代码训练:利用 Gymnasium 注册表
如果你的环境已注册进 Gymnasium(例如官方内置的CartPole-v1)且策略已注册(即使用"MlpPolicy"等内置别名),那么整个训练可以压缩成一行:
from stable_baselines3 import A2C model = A2C("MlpPolicy", "CartPole-v1").learn(10_000)这个"一行训练"之所以可行,源于BaseAlgorithm.__init__中的maybe_make_env()(stable_baselines3/common/base_class.py):当env参数是字符串时,它会自动调用gym.make(env_id, render_mode="rgb_array")创建环境(若环境不支持该参数则回退为不带参数的gym.make),随后照常完成 Monitor / DummyVecEnv 包装与空间校验。注意verbose默认为 0,因此这行代码训练时不会有控制台输出。
从快速示例走向正式项目
docs/guide/quickstart.md还提示:训练中打印的日志输出及字段含义,参见文档 docs/common/logger.md(例如train/policy_loss、train/value_loss、train/explained_variance等指标,这些在 A2C 的train()中通过self.logger.record(...)写入,见 stable_baselines3/a2c/a2c.py)。
当你熟悉了这个最小闭环后,可以沿着官方文档继续深入:
- 训练效果不佳时,参考 docs/guide/rl_tips.md 的调参建议,以及 docs/guide/rl.md 的强化学习基础;
- 接入自定义环境,阅读 docs/guide/custom_env.md;
- 在训练中插入评估、保存、学习率调度等逻辑,阅读 docs/guide/callbacks.md;
- 想要更复杂的观测(图像、字典观测),对应使用
"CnnPolicy"与"MultiInputPolicy",并参考 docs/guide/custom_policy.md; - 尝试其他算法时无需学习新接口——把
A2C换成PPO、DQN、SAC、TD3或DDPG(统一导出见 stable_baselines3/init.py),同样的("MlpPolicy", env)构造方式与learn()/predict()调用即可直接复用。
SB3 快速上手的核心就一句话:sklearn 风格接口 + VecEnv 内部驱动。理解这两点,你就能在几分钟内跑通任意一个受支持算法与环境的训练、推理与可视化全流程。
【免费下载链接】stable-baselines3PyTorch version of Stable Baselines, reliable implementations of reinforcement learning algorithms.项目地址: https://gitcode.com/GitHub_Trending/st/stable-baselines3
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考