Stable-Baselines3 A2C 算法指南:同步优势演员-评论家的原理、参数与工程实践
【免费下载链接】stable-baselines3PyTorch version of Stable Baselines, reliable implementations of reinforcement learning algorithms.项目地址: https://gitcode.com/GitHub_Trending/st/stable-baselines3
A2C(Advantage Actor Critic)是 Stable-Baselines3 中实现的最经典的 on-policy 强化学习算法之一,作为异步 A3C 的同步、确定性变体,它用多环境并行替代了回放缓冲区(replay buffer)。本文以 docs/modules/a2c.md 为主线,结合 stable_baselines3/a2c/a2c.py 等源码实现,系统讲解 A2C 的原理、受支持的环境空间、完整训练/保存/加载/推理代码、关键超参数默认值、CPU 并行优化、训练稳定性技巧(RMSpropTFLike)以及 gSDE 推理注意事项,帮助你在实际项目中正确配置并稳定训练 A2C 智能体。
算法定位:A3C 的同步确定性变体
A2C 是 Asynchronous Advantage Actor Critic (A3C) 论文所提出方法的一个变体。与 A3C 使用多个异步 worker 各自维护一份策略并周期性同步梯度的做法不同,A2C 采用**同步(synchronous)**方式:多个并行环境各自推进n_steps步,收集完整的一批轨迹后,统一更新一次策略网络和价值网络。
由于 A2C 是 on-policy 算法,它不使用回放缓冲区(replay buffer),而是通过**多个并行环境(multiple workers)**来获取多样化的样本,从而在一定程度上缓解样本相关性并稳定训练。这一设计思想在源码中也有直接体现:A2C继承自 stable_baselines3/common/on_policy_algorithm.py 中的OnPolicyAlgorithm,其基类构造函数通过support_multi_env=True明确声明了对多环境并行的原生支持。
从源码结构看,A2C 的类继承关系为:A2C→OnPolicyAlgorithm→BaseAlgorithm,训练数据流为"并行环境收集 rollout → 存入 RolloutBuffer → 一次性完成梯度更新",具体可参考 stable_baselines3/a2c/a2c.py 与 stable_baselines3/common/buffers.py 中的RolloutBuffer实现。
功能支持情况
能力清单
| 能力 | 支持情况 |
|---|---|
| 循环策略(Recurrent policies) | ❌ |
| 多进程(Multi processing) | ✔️ |
Gym 空间支持
| Space | Action(动作空间) | Observation(观测空间) |
|---|---|---|
| Discrete | ✔️ | ✔️ |
| Box | ✔️ | ✔️ |
| MultiDiscrete | ✔️ | ✔️ |
| MultiBinary | ✔️ | ✔️ |
| Dict | ❌ | ✔️ |
这一支持范围与源码中supported_action_spaces的定义完全一致。在 stable_baselines3/a2c/a2c.py 中,A2C.__init__将支持的动作空间限定为spaces.Box、spaces.Discrete、spaces.MultiDiscrete、spaces.MultiBinary四种;而观测空间方面,虽然不支持 Dict 动作空间,但通过MultiInputActorCriticPolicy支持 Dict 形式的观测输入(如表征多个异构输入的 dict 观测)。
快速上手:训练并推理一个 A2C 智能体
以下示例演示在CartPole-v1环境上使用 4 个并行环境训练一个 A2C 智能体,并展示保存、加载与推理的完整流程。注意:该示例仅用于演示库的使用方法,训练得到的智能体不保证能完美解决环境;经过调优的超参数可参考 RL Zoo 仓库。
from stable_baselines3 import A2C from stable_baselines3.common.env_util import make_vec_env # 并行环境:创建 4 个 CartPole-v1 副本 vec_env = make_vec_env("CartPole-v1", n_envs=4) model = A2C("MlpPolicy", vec_env, verbose=1) model.learn(total_timesteps=25000) model.save("a2c_cartpole") del model # 删除模型以演示保存与加载 model = A2C.load("a2c_cartpole") obs = vec_env.reset() while True: action, _states = model.predict(obs) obs, rewards, dones, info = vec_env.step(action) vec_env.render("human")代码要点说明:
make_vec_env("CartPole-v1", n_envs=4)来自 stable_baselines3/common/env_util.py,会创建 4 个并行的向量化环境;A2C("MlpPolicy", vec_env, verbose=1)中verbose=1表示打印设备、所用 wrapper 等基本信息(verbose=2输出 debug 级信息,0为静默);model.learn(total_timesteps=25000)中total_timesteps为训练总步数,log_interval默认 100 表示每 100 次迭代记录一次日志;A2C.load/model.save支持完整的模型序列化(策略网络、优化器状态、超参数等),详见 docs/guide/save_format.md。
优先在 CPU 上运行:并行加速建议
A2C 主要设计为在 CPU 上运行,尤其是未使用 CNN 策略时。为了提升 CPU 利用率,建议关闭 GPU 并使用SubprocVecEnv替代默认的DummyVecEnv:
from stable_baselines3 import A2C from stable_baselines3.common.env_util import make_vec_env from stable_baselines3.common.vec_env import SubprocVecEnv if __name__ == "__main__": env = make_vec_env("CartPole-v1", n_envs=8, vec_env_cls=SubprocVecEnv) model = A2C("MlpPolicy", env, device="cpu") model.learn(total_timesteps=25_000)此处有两点值得注意:
device="cpu"显式指定设备。如果模型被创建在 GPU 上而使用的又是 MlpPolicy,OnPolicyAlgorithm._setup_model会调用_maybe_recommend_cpu()(见 stable_baselines3/common/on_policy_algorithm.py)并发出UserWarning,提示"A2C/PPO 使用 MlpPolicy 时应优先在 CPU 上运行,否则 GPU 利用率低,训练可能更慢"。if __name__ == "__main__"保护。SubprocVecEnv通过多进程创建子环境,在 Windows 等平台上若不放在该保护块内会导致递归进程创建问题;这也是 stable_baselines3/common/vec_env/subproc_vec_env.py 的常规用法要求。
更多关于向量化环境的知识参见 Vectorized Environments 指南。
训练不稳定?换用 RMSpropTFLike 优化器
::: warning 如果发现训练不稳定,或希望复现 stable-baselines(TF 版本)中 A2C 的性能,建议使用stable_baselines3.common.sb2_compat.rmsprop_tf_like中的RMSpropTFLike优化器。可以通过policy_kwargs替换优化器: :::
from stable_baselines3 import A2C from stable_baselines3.common.sb2_compat.rmsprop_tf_like import RMSpropTFLike model = A2C( "MlpPolicy", env, policy_kwargs=dict( optimizer_class=RMSpropTFLike, optimizer_kwargs=dict(eps=1e-5), ), )为什么会有这个建议?
PyTorch 自带的torch.optim.RMSprop与 TensorFlow 原版 RMSProp 存在实现细节差异,而这会影响 A2C 这类对优化器行为敏感的算法。查看 stable_baselines3/common/sb2_compat/rmsprop_tf_like.py 的实现,可以发现RMSpropTFLike相对 PyTorch 原生 RMSprop 做了两处关键修改:
- 把 epsilon 移进平方根内部:更新公式变为
α / (√(v) + ε)中的平方根先加 epsilon,即avg = square_avg.add(eps).sqrt_(),与 TensorFlow 的行为对齐; - 将平方梯度(square_avg)初始化为 1 而不是 0:源码第 103 行
state["square_avg"] = torch.ones_like(p, ...),避免初始步长过大导致早期训练震荡。
默认优化器解析
在 stable_baselines3/a2c/a2c.py 中,A2C.__init__在use_rms_prop=True且用户未显式传入optimizer_class时,会自动将优化器设置为 PyTorch 的th.optim.RMSprop,并附带参数alpha=0.99, eps=rms_prop_eps, weight_decay=0:
if use_rms_prop and "optimizer_class" not in self.policy_kwargs: self.policy_kwargs["optimizer_class"] = th.optim.RMSprop self.policy_kwargs["optimizer_kwargs"] = dict(alpha=0.99, eps=rms_prop_eps, weight_decay=0)也就是说,默认优化器是 RMSprop;若希望改为 Adam,可设置use_rms_prop=False或通过policy_kwargs显式传入optimizer_class。
gSDE 推理注意事项
当使用use_sde=True训练 A2C 模型时(即采用 Generalized State-Dependent Exploration,广义状态依赖探索),需要注意推理阶段的噪声行为:
- 训练过程中,噪声矩阵会在
sde_sample_freq控制的间隔自动重置(sde_sample_freq=-1时仅在 rollout 开始时采样一次); - 但使用
model.predict()进行推理时不会发生自动噪声重置,这会导致即使在deterministic=False的情况下模型行为也趋于确定; - 对于连续控制任务,推荐在推理时使用确定性行为(
deterministic=True); - 如果确实需要推理时的随机行为,必须根据期望的
sde_sample_freq间隔手动调用model.policy.reset_noise(env.num_envs)来重置噪声。
这一行为在OnPolicyAlgorithm.collect_rollouts(见 stable_baselines3/common/on_policy_algorithm.py)中可以看到训练侧的实现:当use_sde且sde_sample_freq > 0、步数满足n_steps % sde_sample_freq == 0时调用self.policy.reset_noise(env.num_envs);而 predict 路径并不包含该逻辑。
核心参数详解
A2C构造函数完整参数及默认值如下(对应 stable_baselines3/a2c/a2c.py):
| 参数 | 默认值 | 含义 |
|---|---|---|
policy | 必填 | 策略模型,如"MlpPolicy"、"CnnPolicy"、"MultiInputPolicy" |
env | 必填 | 学习环境(若已注册到 Gym,可传字符串 ID) |
learning_rate | 7e-4 | 学习率,可以是当前进度剩余比例(1→0)的函数 |
n_steps | 5 | 每个环境每次更新前运行步数(批大小为n_steps × n_env) |
gamma | 0.99 | 折扣因子 |
gae_lambda | 1.0 | GAE 偏差-方差权衡因子,设为 1 时等价于经典 advantage |
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 时每 n 步采样一次新噪声矩阵;-1 表示仅在 rollout 开始时采样 |
rollout_buffer_class | None | 使用的 rollout 缓冲区类,为None时自动选择(Dict 观测用DictRolloutBuffer,否则用RolloutBuffer) |
rollout_buffer_kwargs | None | 创建 rollout 缓冲区时的关键字参数 |
normalize_advantage | False | 是否对 advantage 做归一化 |
stats_window_size | 100 | 用于日志统计的窗口大小,即平均多少个 episode 来报告成功率、平均回合长度与平均回报 |
tensorboard_log | None | TensorBoard 日志目录(None表示不记录) |
policy_kwargs | None | 传给策略的额外参数(如net_arch、optimizer_class),见 A2C Policies 一节 |
verbose | 0 | 日志级别:0 静默、1 信息、2 调试 |
seed | None | 伪随机数生成器种子 |
device | "auto" | 运行设备(cpu/cuda/auto),auto 时优先 GPU |
参数要点解读
n_steps决定批大小:A2C 每次梯度更新使用n_steps × n_env条样本。例如n_envs=8、n_steps=5时,每轮更新批大小为 40。这一参数在OnPolicyAlgorithm.collect_rollouts中作为n_rollout_steps传入,决定单轮 rollout 的长度。normalize_advantage:这是 A2C 相对原始论文实现的一个可选增强。在A2C.train()(见 stable_baselines3/a2c/a2c.py)中,当该参数为True时会对 advantage 做标准归一化:(advantages - advantages.mean()) / (advantages.std() + 1e-8)。learning_rate支持调度函数:可以传入lambda progress_remaining: ...形式的学习率调度器,训练过程中_update_learning_rate会随进度衰减学习率。
源码级解析:A2C 的训练循环与损失计算
A2C 的训练流程可分为"采集"与"更新"两个阶段,理解这两个阶段能帮助你更好地调参。
第一阶段:采集 rollout(collect_rollouts)
在OnPolicyAlgorithm.collect_rollouts中(stable_baselines3/common/on_policy_algorithm.py):
- 策略切换为 eval 模式(
set_training_mode(False)); - 若启用 gSDE,在 rollout 开始(及按
sde_sample_freq)重置噪声; - 循环
n_steps步:以no_grad方式前向计算actions, values, log_probs,执行环境步进,将(obs, actions, rewards, episode_starts, values, log_probs)存入 RolloutBuffer; - 对
TimeLimit.truncated的回合(环境被截断而非真正结束)用价值函数做 bootstrap,避免低估长期回报(对应 GitHub issue #633 的处理); - 最后一帧用价值网络估计
last_values,调用rollout_buffer.compute_returns_and_advantage计算 GAE advantage 与 TD(λ) 回报。
RolloutBuffer.compute_returns_and_advantage(stable_baselines3/common/buffers.py)从后向前递推计算:
delta = rewards[step] + gamma * next_values * next_non_terminal - values[step] last_gae_lam = delta + gamma * gae_lambda * next_non_terminal * last_gae_lam returns = advantages + values当gae_lambda=1.0时,advantage 退化为带价值 bootstrap 的 Monte-Carlo 形式R - V(s),这也是 A2C 默认gae_lambda=1.0的原因——A2C 不使用 GAE 的降方差技巧,而是使用经典 advantage。
第二阶段:单次梯度更新(train)
A2C.train()(stable_baselines3/a2c/a2c.py)每次只对整批数据做一次梯度更新(self.rollout_buffer.get(batch_size=None)会一次性取出全部数据):
- 策略切换为训练模式,更新优化器学习率;
- 计算三项损失:
- 策略梯度损失:
policy_loss = -(advantages * log_prob).mean() - 价值损失:
value_loss = F.mse_loss(rollout_data.returns, values) - 熵正则损失:
entropy_loss = -mean(entropy)(无解析熵时用-mean(-log_prob)近似) - 总损失:
loss = policy_loss + ent_coef * entropy_loss + vf_coef * value_loss
- 策略梯度损失:
- 反向传播后用
clip_grad_norm_以max_grad_norm裁剪梯度,然后优化器更新; - 记录
train/n_updates、train/explained_variance、train/entropy_loss、train/policy_loss、train/value_loss,连续动作下还会记录train/std(对数标准差指数化后的均值)。
从上述实现可以看到ent_coef与vf_coef分别控制探索鼓励与价值学习的权重,max_grad_norm通过梯度裁剪增强稳定性,而normalize_advantage可让不同量纲环境的优势值保持稳定尺度。
策略类型:A2C Policies
A2C通过policy_aliases注册了三种内置策略别名(见 stable_baselines3/a2c/a2c.py 与 stable_baselines3/a2c/policies.py):
| 别名 | 底层类 | 适用场景 |
|---|---|---|
MlpPolicy | ActorCriticPolicy(见 stable_baselines3/common/policies.py) | 低维向量观测(如 CartPole) |
CnnPolicy | ActorCriticCnnPolicy | 图像类观测(如 Atari 游戏) |
MultiInputPolicy | MultiInputActorCriticPolicy | Dict 形式的混合观测(如图像 + 向量) |
它们共享同一个ActorCriticPolicy基类,因此在policy_kwargs中可配置:
net_arch:隐藏层结构,如dict(pi=[64, 64], vf=[64, 64])分别指定策略网络与价值网络,或[64, 64]共享特征提取层;activation_fn:激活函数,默认th.nn.Tanh;optimizer_class/optimizer_kwargs:自定义优化器(如前述RMSpropTFLike);features_extractor_class/features_extractor_kwargs:自定义特征提取器(CNN 场景常用)。
完整参数列表可参考 stable_baselines3/common/policies.py 中ActorCriticPolicy的文档字符串。
实验结果:PyBullet 基准
Atari 游戏
A2C 在 Atari 游戏上的完整学习曲线可在关联的 PR #110 中查看,本文不再展开。
PyBullet 环境
下表为 PyBullet 基准上 2M 步、6 个随机种子的评测结果,其中Gaussian表示使用非结构化高斯噪声探索,gSDE表示使用广义状态依赖探索。超参数取自 gSDE 论文(其超参针对 PyBullet 环境调优):
| Environments | A2C | A2C | PPO | PPO |
|---|---|---|---|---|
| Gaussian | gSDE | Gaussian | gSDE | |
| HalfCheetah | 2003 ± 54 | 2032 ± 122 | 1976 ± 479 | 2826 ± 45 |
| Ant | 2286 ± 72 | 2443 ± 89 | 2364 ± 120 | 2782 ± 76 |
| Hopper | 1627 ± 158 | 1561 ± 220 | 1567 ± 339 | 2512 ± 21 |
| Walker2D | 577 ± 65 | 839 ± 56 | 1230 ± 147 | 2019 ± 64 |
从表中可以观察到:在多数 PyBullet 连续控制任务上,gSDE 探索相比 Gaussian 噪声通常能带来更稳定或更优的表现,且 A2C 的方差普遍小于 PPO 的 Gaussian 配置;但整体上 PPO(尤其是 gSDE 配置)在这些任务上通常能取得更高的绝对分数。完整学习曲线见关联 issue #48。
如何复现实验结果
克隆 rl-zoo 仓库并运行基准测试:
git clone https://github.com/DLR-RM/rl-baselines3-zoo cd rl-baselines3-zoo/运行基准测试(将$ENV_ID替换为上述环境 ID,例如HalfCheetahBulletEnv-v0):
python train.py --algo a2c --env $ENV_ID --eval-episodes 10 --eval-freq 10000绘制结果(此处仅针对 PyBullet 环境):
python scripts/all_plots.py -a a2c -e HalfCheetah Ant Hopper Walker2D -f logs/ -o logs/a2c_results python scripts/plot_from_file.py -i logs/a2c_results.pkl -latex -l A2C参考与延伸阅读
- 原始论文:Asynchronous Methods for Deep Reinforcement Learning(arxiv 1602.01783)
- OpenAI 博客:OpenAI Baselines: ACKTR & A2C
- 向量化环境:见 docs/guide/vec_envs.md
- 本文档为 docs/modules/a2c.md,配套的
A2C自动生成 API 文档由 stable_baselines3/a2c/a2c.py 与 stable_baselines3/common/policies.py 的 docstring 生成 - 相关测试用例可参考 tests/test_run.py、tests/test_sde.py、tests/test_predict.py,覆盖了 A2C 训练、gSDE 与推理行为等核心路径
【免费下载链接】stable-baselines3PyTorch version of Stable Baselines, reliable implementations of reinforcement learning algorithms.项目地址: https://gitcode.com/GitHub_Trending/st/stable-baselines3
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考