
Stable-Baselines3 A2C 算法指南同步优势演员-评论家的原理、参数与工程实践【免费下载链接】stable-baselines3PyTorch version of Stable Baselines, reliable implementations of reinforcement learning algorithms.项目地址: https://gitcode.com/GitHub_Trending/st/stable-baselines3A2CAdvantage 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_envTrue明确声明了对多环境并行的原生支持。从源码结构看A2C 的类继承关系为A2C→OnPolicyAlgorithm→BaseAlgorithm训练数据流为并行环境收集 rollout → 存入 RolloutBuffer → 一次性完成梯度更新具体可参考 stable_baselines3/a2c/a2c.py 与 stable_baselines3/common/buffers.py 中的RolloutBuffer实现。功能支持情况能力清单能力支持情况循环策略Recurrent policies❌多进程Multi processing✔️Gym 空间支持SpaceAction动作空间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_envs4) model A2C(MlpPolicy, vec_env, verbose1) model.learn(total_timesteps25000) 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_envs4)来自 stable_baselines3/common/env_util.py会创建 4 个并行的向量化环境A2C(MlpPolicy, vec_env, verbose1)中verbose1表示打印设备、所用 wrapper 等基本信息verbose2输出 debug 级信息0为静默model.learn(total_timesteps25000)中total_timesteps为训练总步数log_interval默认 100 表示每 100 次迭代记录一次日志A2C.load/model.save支持完整的模型序列化策略网络、优化器状态、超参数等详见 docs/guide/save_format.md。优先在 CPU 上运行并行加速建议A2C 主要设计为在 CPU 上运行尤其是未使用 CNN 策略时。为了提升 CPU 利用率建议关闭 GPU 并使用SubprocVecEnv替代默认的DummyVecEnvfrom 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_envs8, vec_env_clsSubprocVecEnv) model A2C(MlpPolicy, env, devicecpu) model.learn(total_timesteps25_000)此处有两点值得注意devicecpu显式指定设备。如果模型被创建在 GPU 上而使用的又是 MlpPolicyOnPolicyAlgorithm._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-baselinesTF 版本中 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_kwargsdict( optimizer_classRMSpropTFLike, optimizer_kwargsdict(eps1e-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_propTrue且用户未显式传入optimizer_class时会自动将优化器设置为 PyTorch 的th.optim.RMSprop并附带参数alpha0.99, epsrms_prop_eps, weight_decay0if 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(alpha0.99, epsrms_prop_eps, weight_decay0)也就是说默认优化器是 RMSprop若希望改为 Adam可设置use_rms_propFalse或通过policy_kwargs显式传入optimizer_class。gSDE 推理注意事项当使用use_sdeTrue训练 A2C 模型时即采用 Generalized State-Dependent Exploration广义状态依赖探索需要注意推理阶段的噪声行为训练过程中噪声矩阵会在sde_sample_freq控制的间隔自动重置sde_sample_freq-1时仅在 rollout 开始时采样一次但使用model.predict()进行推理时不会发生自动噪声重置这会导致即使在deterministicFalse的情况下模型行为也趋于确定对于连续控制任务推荐在推理时使用确定性行为deterministicTrue如果确实需要推理时的随机行为必须根据期望的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、MultiInputPolicyenv必填学习环境若已注册到 Gym可传字符串 IDlearning_rate7e-4学习率可以是当前进度剩余比例1→0的函数n_steps5每个环境每次更新前运行步数批大小为n_steps × n_envgamma0.99折扣因子gae_lambda1.0GAE 偏差-方差权衡因子设为 1 时等价于经典 advantageent_coef0.0损失计算中的熵系数vf_coef0.5损失计算中的价值函数系数max_grad_norm0.5梯度裁剪最大值rms_prop_eps1e-5RMSProp 的 epsilon稳定分母平方根计算use_rms_propTrue是否使用 RMSprop默认而非 Adamuse_sdeFalse是否使用 gSDE 替代动作噪声探索sde_sample_freq-1使用 gSDE 时每 n 步采样一次新噪声矩阵-1 表示仅在 rollout 开始时采样rollout_buffer_classNone使用的 rollout 缓冲区类为None时自动选择Dict 观测用DictRolloutBuffer否则用RolloutBufferrollout_buffer_kwargsNone创建 rollout 缓冲区时的关键字参数normalize_advantageFalse是否对 advantage 做归一化stats_window_size100用于日志统计的窗口大小即平均多少个 episode 来报告成功率、平均回合长度与平均回报tensorboard_logNoneTensorBoard 日志目录None表示不记录policy_kwargsNone传给策略的额外参数如net_arch、optimizer_class见 A2C Policies 一节verbose0日志级别0 静默、1 信息、2 调试seedNone伪随机数生成器种子deviceauto运行设备cpu/cuda/autoauto 时优先 GPU参数要点解读n_steps决定批大小A2C 每次梯度更新使用n_steps × n_env条样本。例如n_envs8、n_steps5时每轮更新批大小为 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 的训练流程可分为采集与更新两个阶段理解这两个阶段能帮助你更好地调参。第一阶段采集 rolloutcollect_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_advantagestable_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_lambda1.0时advantage 退化为带价值 bootstrap 的 Monte-Carlo 形式R - V(s)这也是 A2C 默认gae_lambda1.0的原因——A2C 不使用 GAE 的降方差技巧而是使用经典 advantage。第二阶段单次梯度更新trainA2C.train()stable_baselines3/a2c/a2c.py每次只对整批数据做一次梯度更新self.rollout_buffer.get(batch_sizeNone)会一次性取出全部数据策略切换为训练模式更新优化器学习率计算三项损失策略梯度损失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 PoliciesA2C通过policy_aliases注册了三种内置策略别名见 stable_baselines3/a2c/a2c.py 与 stable_baselines3/a2c/policies.py别名底层类适用场景MlpPolicyActorCriticPolicy见 stable_baselines3/common/policies.py低维向量观测如 CartPoleCnnPolicyActorCriticCnnPolicy图像类观测如 Atari 游戏MultiInputPolicyMultiInputActorCriticPolicyDict 形式的混合观测如图像 向量它们共享同一个ActorCriticPolicy基类因此在policy_kwargs中可配置net_arch隐藏层结构如dict(pi[64, 64], vf[64, 64])分别指定策略网络与价值网络或[64, 64]共享特征提取层activation_fn激活函数默认th.nn.Tanhoptimizer_class/optimizer_kwargs自定义优化器如前述RMSpropTFLikefeatures_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 环境调优EnvironmentsA2CA2CPPOPPOGaussiangSDEGaussiangSDEHalfCheetah2003 ± 542032 ± 1221976 ± 4792826 ± 45Ant2286 ± 722443 ± 892364 ± 1202782 ± 76Hopper1627 ± 1581561 ± 2201567 ± 3392512 ± 21Walker2D577 ± 65839 ± 561230 ± 1472019 ± 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-v0python 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 Learningarxiv 1602.01783OpenAI 博客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),仅供参考