ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

Stable-Baselines3 DDPG 算法实战指南:连续动作空间的确定性策略梯度实现

Stable-Baselines3 DDPG 算法实战指南:连续动作空间的确定性策略梯度实现 Stable-Baselines3 DDPG 算法实战指南连续动作空间的确定性策略梯度实现【免费下载链接】stable-baselines3PyTorch version of Stable Baselines, reliable implementations of reinforcement learning algorithms.项目地址: https://gitcode.com/GitHub_Trending/st/stable-baselines3DDPGDeep Deterministic Policy Gradient是 Stable-Baselines3 中面向连续动作空间的经典离策略强化学习算法它将 DQN 的经验回放与目标网络技巧同确定性策略梯度相结合。本文以 docs/modules/ddpg.md 为主线结合仓库源码stable_baselines3/ddpg/ddpg.py、stable_baselines3/td3/td3.py深入讲解其在 SB3 中的实现原理、支持的 Gym 空间、完整示例代码、参数含义与基准复现方法读完即可在 PyTorch 环境下独立完成 DDPG 的训练、保存、加载与推理。DDPG 核心思想DQN 技巧与确定性策略梯度的结合DDPG 由 Lillicrap 等人于 2016 年提出其设计目标是让 DQN 这一类基于值函数的方法能够处理连续动作空间。其核心组合方式如下确定性策略梯度Deterministic Policy Gradient, Silver 2014策略网络直接输出确定性的动作 $a \mu(s)$而非像 REINFORCE / Actor-Critic 那样采样概率分布从而把对动作空间的期望积分转化为可直接求梯度的确定性映射显著降低连续空间中的采样方差DQN 的两大技巧使用**经验回放Replay Buffer打破样本时间相关性使用软更新的目标网络Target Network**稳定 Q 值更新缓解自举带来的发散问题。在 Stable-Baselines3 中DDPG 与 TD3 共享同一套策略网络与算法基座。正如官方文档的 note 所述DDPG 可视为其继任者 TD3 的一个特例二者共用相同的策略与实现代码。这一设计决策直接体现在源码中DDPG类继承自TD3类stable_baselines3/ddpg/ddpg.py。源码视角DDPG 如何作为 TD3 的特例实现查看 stable_baselines3/ddpg/ddpg.py 可以发现DDPG类只做两件事调用父类TD3的构造函数并关闭 TD3 引入的改进技巧超参数TD3 默认值DDPG 传入值说明policy_delay21TD3 每 2 次 critic 更新才更新一次策略DDPG 每步都更新target_policy_noise0.20.1目标策略平滑噪声标准差DDPG 中因target_noise_clip0.0而被钳制为零target_noise_clip0.50.0目标策略噪声裁剪范围置 0 意味着完全不添加平滑噪声n_critics21TD3 使用双 Q 网络取最小值DDPG 只保留单个 critic源码中对应的关键注释stable_baselines3/ddpg/ddpg.py# Remove all tricks from TD3 to obtain DDPG: # we still need to specify target_policy_noise 0 to avoid errors policy_delay1, target_noise_clip0.0, target_policy_noise0.1,这里有一个容易被忽视的实现细节即使 DDPG 用不到目标策略平滑噪声target_policy_noise仍必须传入大于 0 的值以避免内部报错。此外构造函数会在policy_kwargs中强制n_critics 1stable_baselines3/ddpg/ddpg.py因此即使你在policy_kwargs里指定n_critics2DDPG 也会退化为单 critic。由于 DDPG 完全复用 TD3 的训练循环stable_baselines3/td3/td3.py当policy_delay1、target_noise_clip0.0时训练流程自然退化为经典 DDPG从回放缓冲区采样一个 mini-batch用目标 actor 生成目标动作由于噪声裁剪为 0next_actions actor_target(next_obs)用目标 critic 计算目标 Q 值target_q reward (1 - done) * gamma * next_q最小化当前 Q 与目标 Q 的 MSE 更新 critic每个训练步都更新 actor最大化critic.q1_forward(obs, actor(obs))以系数tau对 actor/critic 的目标网络做 Polyak 软更新。测试用例也印证了这一点在 tests/test_run.py 中test_deterministic_pg将TD3与DDPG参数化后使用同一种噪声NormalActionNoise或OrnsteinUhlenbeckActionNoise在Pendulum-v1上验证训练流程。可用策略与网络架构DDPG 与 TD3 共享策略类全部定义在 stable_baselines3/td3/policies.py 中三种策略如下策略观测空间类型特征提取器说明MlpPolicyBox / Discrete 等向量观测FlattenExtractor等价于TD3Policy默认全连接架构CnnPolicy图像观测NatureCNN卷积特征提取 MLP 头MultiInputPolicyDict字典观测CombinedExtractor处理混合类型图像 向量观测网络架构默认值stable_baselines3/td3/policies.py遵循原始论文MLP 策略默认net_arch [400, 300]两层隐藏层400 与 300 个神经元使用NatureCNN时默认net_arch [256, 256]。可通过policy_kwargsdict(net_arch[64, 64])自定义测试用例即采用这一写法。Actor 网络stable_baselines3/td3/policies.py是一个确定性映射输入观测 → 特征提取 → MLP → 输出动作。关键点是squash_outputTrue即输出层经过tanh激活将动作压缩到[-1, 1]区间这与 Gymnasium 的Box动作空间通常需要先做rescale或环境本身接受[-1,1]相匹配。Critic 网络使用ContinuousCritic定义于 stable_baselines3/common/policies.py输入为「观测 动作」拼接输出 Q 值。目标网络actor_target、critic_target在构造时通过load_state_dict复制在线网络权重并始终处于eval模式stable_baselines3/td3/policies.py。支持范围Can I use?依据原文档DDPG 的能力边界如下循环策略Recurrent policies不支持 ❌多进程Multi processing支持 ✔️底层通过OffPolicyAlgorithm的support_multi_envTrue启用见 stable_baselines3/td3/td3.pyGym 空间支持矩阵SpaceActionObservationDiscrete❌✔️Box✔️✔️MultiDiscrete❌✔️MultiBinary❌✔️Dict❌✔️动作空间仅支持Box连续动作这与算法定义一致——确定性策略梯度需要可求梯度的连续动作输出而观测空间则非常灵活向量、离散、图像、字典观测均可处理。因此 DDPG 的典型应用场景是机器人控制、连续动力系统控制等任务。快速上手在 Pendulum-v1 上训练 DDPG原文档给出的示例使用 Gymnasium 的Pendulum-v1环境经典倒立摆连续控制问题完整代码如下import gymnasium as gym import numpy as np from stable_baselines3 import DDPG from stable_baselines3.common.noise import NormalActionNoise, OrnsteinUhlenbeckActionNoise env gym.make(Pendulum-v1, render_modergb_array) # The noise objects for DDPG n_actions env.action_space.shape[-1] action_noise NormalActionNoise(meannp.zeros(n_actions), sigma0.1 * np.ones(n_actions)) model DDPG(MlpPolicy, env, action_noiseaction_noise, verbose1) model.learn(total_timesteps10000, log_interval10) model.save(ddpg_pendulum) vec_env model.get_env() del model # remove to demonstrate saving and loading model DDPG.load(ddpg_pendulum) obs vec_env.reset() while True: action, _states model.predict(obs) obs, rewards, dones, info vec_env.step(action) env.render(human)代码要点说明必须配置action_noiseDDPG 的策略是确定性的若不给动作添加噪声训练初期将因缺乏探索而无法学到有效策略训练规模示例仅 10000 步目的是演示 API 用法训练出的智能体不一定能真正解决环境生产级超参数可在 RL Zoo 仓库rl-baselines3-zoo中找到保存与加载model.save()会保存模型权重与超参数DDPG.load()加载后通过model.get_env()拿回训练时的向量化环境随后model.predict(obs)输出确定性动作DDPG 的预测始终是确定性的deterministic参数被忽略见 stable_baselines3/td3/policies.py。探索噪声详解DDPG 的探索噪声定义在 stable_baselines3/common/noise.py常用两类1. NormalActionNoise高斯噪声——推荐首选NormalActionNoise(meannp.zeros(n_actions), sigma0.1 * np.ones(n_actions))每个时间步从均值为mean、标准差为sigma的高斯分布中独立采样stable_baselines3/common/noise.py即np.random.normal(mu, sigma)。sigma控制探索强度实践中常取0.1或0.2乘以动作维度单位向量。2. OrnsteinUhlenbeckActionNoiseOU 噪声——经典 DDPG 论文所用具有时间相关性近似带摩擦的布朗运动OrnsteinUhlenbeckActionNoise(meannp.zeros(n_actions), sigma0.1 * np.ones(n_actions))其递推公式stable_baselines3/common/noise.py为noise noise_prev theta * (mean - noise_prev) * dt sigma * sqrt(dt) * N(0, 1)默认参数theta0.15、dt1e-2控制均值回归速度。OU 噪声在每一步后更新并保存noise_prev回合结束时调用reset()归位stable_baselines3/common/noise.py。对于大多数任务独立高斯噪声通常已足够且更简单高效。3. 多环境适配若使用向量化多环境训练SB3 会在内部用VectorizedActionNoise为每个环境深拷贝独立的噪声发生器stable_baselines3/common/noise.py保证各环境探索相互独立。DDPG 参数详解DDPG构造函数完整签名位于 stable_baselines3/ddpg/ddpg.py参数含义如下参数默认值含义policy必填策略类名MlpPolicy/CnnPolicy/MultiInputPolicyenv必填Gymnasium 环境实例或注册名字符串learning_rate1e-3Adam 优化器学习率actor 与 critic 共用可传(1, 0)形式的调度函数随训练进度从 1 线性衰减到 0buffer_size1_000_000回放缓冲区容量默认 100 万条经验learning_starts100开始学习前需收集的步数预填充回放缓冲区batch_size256每次梯度更新的 mini-batch 大小tau0.005目标网络软更新系数Polyak update01 之间gamma0.99折扣因子train_freq1更新频率整数每 N 步或元组如(5, step)、(2, episode)gradient_steps1每次 rollout 后的梯度更新次数-1表示与环境交互步数等量action_noiseNone探索噪声DDPG 必须显式指定否则无法有效探索replay_buffer_classNone自定义回放缓冲区类如HerReplayBufferNone时自动选择replay_buffer_kwargsNone传给回放缓冲区的额外关键字参数optimize_memory_usageFalse启用回放缓冲区的内存优化变体以增加复杂度为代价n_steps1大于 1 时使用 n-step 回报配合NStepReplayBuffer更新 Q 网络tensorboard_logNoneTensorBoard 日志目录None则不记录policy_kwargsNone传给策略的额外参数如net_arch、activation_fn、n_critics等verbose0日志级别0 无输出1 输出设备/包装器等信息2 输出调试信息seedNone伪随机数生成器种子deviceauto运行设备auto表示 GPU 可用时自动用 GPUlearn()方法参数total_timesteps训练总步数、callback回调函数、log_interval默认 4每隔 N 次 rollout 记录一次日志、tb_log_name默认DDPGTensorBoard 运行名、reset_num_timesteps默认True、progress_bar是否显示进度条。策略层参数policy_kwargs见 stable_baselines3/td3/policies.pynet_arch默认[400, 300]、activation_fn默认nn.ReLU、features_extractor_class、normalize_images默认True图像观测自动除以 255、optimizer_class默认th.optim.Adam、share_features_extractor默认Falseactor 与 critic 是否共享特征提取器。基准结果与复现原文档给出了 PyBullet 基准上的对比结果1M 步、6 个随机种子其中 DDPG 使用高斯噪声Gaussian探索TD3 使用高斯噪声SAC 使用 gSDE广义状态依赖探索。注意DDPG 采用了 gSDE 论文中 TD3 的超参数。环境DDPG (Gaussian)TD3 (Gaussian)SAC (gSDE)HalfCheetah2272 ± 692774 ± 352984 ± 202Ant1651 ± 4073305 ± 433102 ± 37Hopper1201 ± 2112429 ± 1262262 ± 1Walker2D882 ± 1862063 ± 1852136 ± 67完整学习曲线参见该仓库的 issue #48。从表中可以直观看到 TD3 与 SAC 相对 DDPG 的显著提升——这正是后续算法在缓解函数近似误差上取得进步的证据。复现步骤克隆 RL Zoo 仓库并安装依赖后按以下命令复现将$ENV_ID替换为上表中的环境名如HalfCheetahBulletEnv-v0git clone https://github.com/DLR-RM/rl-baselines3-zoo cd rl-baselines3-zoo/运行基准训练每 10000 步评估 10 个回合python train.py --algo ddpg --env $ENV_ID --eval-episodes 10 --eval-freq 10000绘制结果曲线python scripts/all_plots.py -a ddpg -e HalfCheetah Ant Hopper Walker2D -f logs/ -o logs/ddpg_results python scripts/plot_from_file.py -i logs/ddpg_results.pkl -latex -l DDPG实践要点小结动作空间只能是连续Box离散/多离散/字典动作请改用 DQN 或 PPO 等算法探索噪声不可省略NormalActionNoise优先于 OU 噪声DDPG 是 TD3 的特例单 critic、无目标平滑噪声、无延迟策略更新若任务对超参数敏感可直接换用 TD3TD3 文档以降低调参难度目标网络软更新系数tau不宜过大默认 0.005过大会导致训练不稳定使用CnnPolicy处理图像观测时默认网络架构会切换为[256, 256]且支持Dict观测的MultiInputPolicy便于混合多模态输入。【免费下载链接】stable-baselines3PyTorch version of Stable Baselines, reliable implementations of reinforcement learning algorithms.项目地址: https://gitcode.com/GitHub_Trending/st/stable-baselines3创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表