强化学习Stable-Baselines3


Stable-Baselines3(SB3)是基于 PyTorch 实现的强化学习库,日常使用(直接用 PPO/SAC/TD3 等算法训练)一般不需要你手写 PyTorch 代码;只有在自定义网络/特征提取器/导出模型等场景,才会写到 PyTorch。

项目地址:https://github.com/DLR-RM/stable-baselines3
文档地址:https://stable-baselines3.readthedocs.io/en/master/guide/quickstart.html

pip install stable-baselines3[extra] 完整安装包括Tensorboard、OpenCV 、ale-py
pip install stable-baselines3

算法选型

离散动作(如:开关/档位/网格移动/高层策略选择)

  • DQN(Deep Q-Network)
    • 类型:值函数、离线可复用(off-policy)
    • 特征:仅支持离散动作;经验回放,样本效率较好;对超参和探索策略(ε-greedy)较敏感
    • 何时用:动作可枚举(如 5~20 个选项);需要复用历史数据/离线数据;高层决策器(选路线/工位/模式)

连续动作(如:速度、力矩、关节角速度)

  • DDPG(Deep Deterministic Policy Gradient)

    • 类型:确定性策略、off-policy
    • 特征:早期连续控制基线,对噪声与超参敏感
    • 何时用:纯教学/做对比基线;一般更建议直接选 TD3 或 SAC
  • TD3(Twin Delayed DDPG)

    • 类型:确定性策略、off-policy
    • 特征:双 Q、延迟策略更新、目标策略平滑→比 DDPG 更稳定、抗过估计
    • 何时用:低维连续控制(机械臂关节、移动底盘),想要高样本效率且环境较“干净”;实机上更稳
  • SAC(Soft Actor-Critic)

    • 类型:随机策略、off-policy、熵正则
    • 特征:鲁棒、样本效率高,对超参没那么挑;随机性帮助探索
    • 何时用:机器人连续控制首选(臂/轮/抓取),尤其是需要复用数据(离线+在线混用)、对扰动/噪声更友好
    • 备注:SB3 的 SAC 主要用于连续动作

通用/“大力出奇迹”型(易并行、开箱即用)

  • A2C(Advantage Actor-Critic)

    • 类型:on-policy(不复用旧数据)
    • 特征:实现简单、训练快,样本效率一般;适合 CPU 并行多个环境
    • 何时用:快速原型、小项目或教学;资源有限但想先跑通
  • PPO(Proximal Policy Optimization)

    • 类型:on-policy
    • 特征:稳健、容错好、并行友好(VecEnv/SubprocVecEnv/大批并行仿真);论文/工程两开花
    • 何时用:默认首选通用解(离散/连续皆可);Isaac/MuJoCo 大并行时常作为主力基线
    • 注意:on-policy → 样本效率低于 SAC/TD3,仿真并行能弥补

稀疏奖励/目标导向任务的“外挂”

  • HER(Hindsight Experience Replay)(SB3 作为包装器提供)
    • 类型:与 off-policy 算法(DDPG/TD3/SAC)搭配
    • 特征:把失败轨迹“事后改写目标”,缓解稀疏奖励
    • 何时用:目标导向(例如到达某位置/抓取到某物),奖励稀疏或全 0;常配 SAC/TD3

场景化选型(直接照着挑)

机器人连续控制(机械臂/移动底盘/平衡车/四足)

  • 仿真可大并行(Isaac/MuJoCo):先用 PPO 快速拿到可用策略 → 再用 SAC 做精细/鲁棒版本
  • 实机样本宝贵:SAC > TD3(能复用旧数据/离线数据,样本效率更好)
  • 稀疏奖励/到达目标:HER + SAC/TD3

离散动作或高层决策(模式/流程/工序选择)

  • DQN:动作离散、需要复用历史数据
  • PPO:若仿真可并行、希望调参简单且稳

需要更快原型/教学/入门

  • A2C(极简)或 PPO(更稳,社区资料最多)

像素输入(相机/深度图直接端到端)

  • PPO(CnnPolicy) 或 SAC(CnnPolicy):配合大 batch + 多环境并行;像素任务上 GPU 会更有用

官方教学演示

py
import gymnasium as gym                     # 引入 Gymnasium(新版 Gym),用于创建强化学习环境  # :环境提供观测/奖励/结束信号等接口
from stable_baselines3 import A2C           # 从 Stable-Baselines3 导入 A2C 算法                 # :A2C 是 on-policy 的 actor-critic 算法

# 创建环境
env = gym.make("CartPole-v1",               # 创建 CartPole 平衡杆环境(离散动作,经典教学环境)     # :观测为杆角度/小车位置等
               render_mode="rgb_array")     # 将渲染模式设为 'rgb_array'(返回帧数组而非开窗口)     # :如果想弹窗显示,应使用 render_mode='human'

# 创建模型
model = A2C("MlpPolicy",                    # 选择 MLP 策略(MlpPolicy多层感知机,适合低维观测),CnnPolicy图像 MultiInputPolicy字典观测 三选一,或policy_kwargs自定义
            env,                            # 传入刚刚创建的环境,SB3 会自动封装为向量环境(DummyVecEnv) # :即使只有1个环境也会向量化
            verbose=1,                      # verbose=1 打印训练日志(如 timesteps、loss 等)        # :便于观察训练过程
            device="cpu")                   # 强制把模型放在 CPU 上训练(env环境始终在cpu上运行,而PyTorch 的网络与相关张量的计算放在 CPU 上进行也可以指定gpu)
model.learn(total_timesteps=10_000)         # 训练 10,000 个时间步                                   # :时间步越大,策略一般越稳定(视任务而定)

vec_env = model.get_env()                   # 取出模型内部保存的向量化环境(VecEnv 实例)             # :便于统一处理 reset/step/render 等
obs = vec_env.reset()                       # 重置向量环境,返回初始观测(形状通常为 [n_envs, obs_dim]) # :这里 n_envs=1,因此 obs 形如 (1, obs_dim)
# 评估模型
for i in range(1000):                       # 进行 1000 次评估/测试步(不更新模型参数)               # :只是用训练好的策略跑
    action, _state = model.predict(         # 用训练好的策略推理动作(推理时无需梯度)               # :_state 在RNN策略时用于隐藏状态,这里无用
        obs,                                # 输入当前观测(向量环境下是批量观测)
        deterministic=True)                 # 设为确定性策略:选择均值/概率最大的动作(减少随机性)      # :便于评估表现
    obs, reward, done, info = vec_env.step( # 与环境交互一步:返回新观测、奖励、是否结束、调试信息     # :在 VecEnv 中这些通常是批量(ndarray/list)
        action)                             # 将动作送入环境;VecEnv 会按批次广播到各子环境           # :n_envs=1 时 action 形如 (1,)
    vec_env.render("human")                 # 请求以“human”模式渲染(窗口显示)                       # :但当前 env 是 rgb_array 模式,可能不会弹窗