CleanRL DDPG 基准评测:PyTorch 与 JAX 双实现在 MuJoCo v4 连续控制上的复现、对比与实现细节
CleanRL DDPG 基准评测PyTorch 与 JAX 双实现在 MuJoCo v4 连续控制上的复现、对比与实现细节【免费下载链接】cleanrlHigh-quality single file implementation of Deep Reinforcement Learning algorithms with research-friendly features (PPO, DQN, C51, DDPG, TD3, SAC, PPG)项目地址: https://gitcode.com/GitHub_Trending/cl/cleanrl本文以 CleanRL 仓库中的基准评测文档 docs/benchmark/ddpg.md 为核心结合算法说明文档 docs/rl-algorithms/ddpg.md系统梳理 DDPGDeep Deterministic Policy Gradient在连续控制任务上的官方基准结果并给出可一键复现的命令、日志指标解读、与原始论文的逐项实现差异以及 PyTorch 与 JAX 两种实现的选择依据。读完本文你将能够理解这份评测表的来历在本地跑通单次训练与全量基准并能读懂其背后每一处工程决策。一、评测背景CleanRL 中的 DDPG 双实现DDPG 是深度强化学习中处理连续控制问题的经典算法它将 DQN 扩展到连续动作空间引入一个直接输出连续动作的确定性 Actor同时沿用 DQN 的经验回放replay buffer与目标网络target network两大技巧。CleanRL 在仓库中提供了两份单文件实现实现文件技术栈适用场景cleanrl/ddpg_continuous_action.pyPyTorch连续动作空间低维Box观测cleanrl/ddpg_continuous_action_jax.pyJAX Flax Optax连续动作空间低维Box观测追求吞吐两份实现均以Hopper-v4为默认环境支持 Gymnasium 的Box观测空间与Box连续动作空间并共享同一套命令行参数接口由tyro自动生成。基准评测正是在这两个实现之间展开。二、评测结果总览六个 MuJoCo v4 环境的 PyTorch 与 JAX 对比以下数据来自仓库中的 docs/benchmark/ddpg.md为 3 个随机种子下的平均回合回报mean ± std评测环境为 Gymnasium MuJoCo v4环境ddpg_continuous_actionPyTorchddpg_continuous_action_jaxJAXHalfCheetah-v410374.07 ± 157.378638.60 ± 1954.46Walker2d-v41240.16 ± 390.101427.23 ± 104.91Hopper-v41576.78 ± 818.981208.52 ± 659.22InvertedPendulum-v4642.68 ± 69.56804.30 ± 87.60Humanoid-v41699.56 ± 694.221513.61 ± 248.60Pusher-v4-77.30 ± 38.78-38.56 ± 4.47从结果看两个实现在多数环境上量级相当PyTorch 版在 HalfCheetah-v4、Hopper-v4、Humanoid-v4 上略高而 JAX 版在 Walker2d-v4、InvertedPendulum-v4、Pusher-v4 上表现更稳标准差更小。需要强调的是这份数据主要用于验证实现质量与对比两套后端的数值行为而非在不同超参搜索下追求某一环境的绝对最优——同一环境上两个实现的差距更多反映的是后端浮点行为、随机种子与训练动态的差异。三、单次训练在本地复现一份 DDPG3.1 环境安装PyTorch 版只需 MuJoCo 相关依赖requirements/requirements-mujoco.txtJAX 版还需追加 JAX 依赖requirements/requirements-jax.txt。使用 uv 或 pip 二选一# 方式一uv仓库使用 uv 管理依赖 uv pip install . uv run python cleanrl/ddpg_continuous_action.py --help uv pip install .[mujoco] # JAX 版额外安装 jax extra uv pip install .[mujoco, jax]# 方式二pip pip install -r requirements/requirements-mujoco.txt python cleanrl/ddpg_continuous_action.py --help python cleanrl/ddpg_continuous_action.py --env-id Hopper-v4JAX 版对应命令为pip install -r requirements/requirements-jax.txt python cleanrl/ddpg_continuous_action_jax.py --env-id Hopper-v4注意JAX 官方不支持 Windows 原生环境Windows 用户请通过 WSLWindows Subsystem for Linux安装运行。3.2 核心命令行参数两份实现共享同一套参数见源码中的Argsdataclass常用参数及默认值如下参数默认值含义--env-idHopper-v4环境 ID--total-timesteps1000000总训练步数--learning-rate3e-4优化器学习率--buffer-size1000000经验回放容量--gamma0.99折扣因子--tau0.005目标网络软更新系数--batch-size256采样 batch 大小--exploration-noise0.1探索噪声尺度--learning-starts25000开始学习前收集的步数--policy-frequency2策略Actor更新频率实现延迟更新--seed1随机种子--torch-deterministicTrue是否启用 cuDNN 确定性--trackFalse是否用 Weights Biases 跟踪实验--capture-videoFalse是否录制智能体视频到videos/--save-modelFalse训练结束后是否保存模型到runs/{run_name}/--upload-modelFalse是否将模型上传至 Hugging Face Hub在learning_starts之前代码从动作空间均匀采样envs.single_action_space.sample()进行纯随机探索之后才切换到 Actor 输出并叠加高斯噪声噪声尺度为actor.action_scale * args.exploration_noise最后裁剪回动作空间边界。四、批量运行官方基准benchmark/ddpg.sh 全解读官方基准脚本 benchmark/ddpg.sh 用 CleanRL 自带的cleanrl_utils.benchmark模块在 6 个环境上各跑 3 个种子。PyTorch 版命令如下uv pip install .[mujoco] python -m cleanrl_utils.benchmark \ --env-ids HalfCheetah-v4 Walker2d-v4 Hopper-v4 InvertedPendulum-v4 Humanoid-v4 Pusher-v4 \ --command uv run python cleanrl/ddpg_continuous_action.py --track \ --num-seeds 3 \ --workers 18 \ --slurm-gpus-per-task 1 \ --slurm-ntasks 1 \ --slurm-total-cpus 10 \ --slurm-template-path benchmark/cleanrl_1gpu.slurm_templateJAX 版在安装上多一步脚本中同时安装.[mujoco, jax]并升级到 CUDA 版 JAX命令本身结构与上相同仅将执行目标替换为cleanrl/ddpg_continuous_action_jax.pyuv pip install .[mujoco, jax] uv pip install --upgrade jax[cuda11_cudnn82]0.4.8 -f https://storage.googleapis.com/jax-releases/jax_cuda_releases.html uv run python -m cleanrl_utils.benchmark \ --env-ids HalfCheetah-v4 Walker2d-v4 Hopper-v4 InvertedPendulum-v4 Humanoid-v4 Pusher-v4 \ --command uv run python cleanrl/ddpg_continuous_action_jax.py --track \ --num-seeds 3 \ --workers 18 \ --slurm-gpus-per-task 1 \ --slurm-ntasks 1 \ --slurm-total-cpus 10 \ --slurm-template-path benchmark/cleanrl_1gpu.slurm_template参数要点--num-seeds 3每个环境跑 3 个随机种子评测表里的mean ± std即来源于此--workers 18并发调度的任务数配合 SLURM 集群资源每任务 1 块 GPU、10 个 CPU批量提交--track训练时用 wandb 记录charts/episodic_return这是后续绘制学习曲线ddpg_plot.sh的数据来源若在单机非 SLURM环境可去掉--slurm-*相关参数直接串行/并行运行。五、日志指标解读TensorBoard 里每一行都在说什么运行后 TensorBoard 会自动记录以下指标对应源码中的writer.add_scalar调用charts/episodic_return每个 episode 的回合回报用于直观判断收敛水平charts/SPS每秒环境步数Steps Per Second用于衡量训练吞吐losses/qf1_lossQ 网络的一步时序差分均方误差。形式化地Critic 目标为$$ J(\theta^{Q}) \mathbb{E}_{(s,a,r,s) \sim \mathcal{D}} \big[ (Q(s, a) - y)^2 \big], $$其中 Bellman 更新目标 $y r \gamma Q(s, a)$$a \sim \mu(s)$ 由目标 Actor 给出$\mathcal{D}$ 为回放缓冲区。在源码中对应 cleanrl/ddpg_continuous_action.py 的qf1_loss F.mse_loss(qf1_a_values, next_q_value)losses/actor_loss实现为-qf1(data.observations, actor(data.observations)).mean()即基于当前观测与 Actor 输出动作计算的负平均 Q 值。最小化该损失等价于沿确定性策略梯度更新 Actor$$ \nabla_{\theta^{\mu}} J \approx \frac{1}{N}\sum_i\left.\left.\nabla_{a} Q\left(s, a \mid \theta^{Q}\right)\right|{ss{i}, a\mu\left(s_{i}\right)} \nabla_{\theta^{\mu}} \mu\left(s \mid \theta^{\mu}\right)\right|{s{i}} $$losses/qf1_values实现为qf1(data.observations, data.actions).view(-1)的均值即回放样本的平均 Q 值可用于判断 Q 值的低估/高估倾向。六、实现细节CleanRL DDPG 与原始论文的差异CleanRL 的 PyTorch 实现参考了sfujim/TD3仓库的OurDDPG.py因此与 Lillicrap et al. (2016) 原始论文存在多处工程差异以下逐条均为文档与源码共同确认的事实探索噪声CleanRL 使用高斯噪声 $\mathcal{N}(0, 0.1)$而原始论文使用 Ornstein-Uhlenbeck 过程$\theta0.15, \sigma0.2$。实践中高斯噪声更简单且往往足够这是后续 DDPG 类实现的主流选择环境CleanRL 基于开源 Gymnasium MuJoCo 环境原始论文使用其内部的专有 MuJoCo 环境网络结构CleanRL 的 Actor/Critic 隐藏层均为 256 维fc1256, fc2256原始论文为 400/300 且 Critic 在第二层拼接动作。CleanRL 的 Critic 在输入层就直接拼接[obs, action]见 cleanrl/ddpg_continuous_action.py 中QNetwork.forward的torch.cat([x, a], 1)学习率CleanRL 对 Actor 与 Critic 统一使用3e-4原始论文分别为1e-3与1e-4batch 与 tauCleanRL 使用--batch-size 256 --tau 0.005原始论文为--batch-size 64 --tau 0.001非对称动作空间支持原始实现假定动作区间为 $[-1,1]$用max_action * tanh(...)缩放而 CleanRL 通过action_scale与action_bias两个 buffer 支持任意区间乃至非对称区间的动作空间# 动作重缩放来自 Actor 的注册 buffer self.register_buffer( action_scale, torch.tensor((env.action_space.high - env.action_space.low) / 2.0, dtypetorch.float32) ) self.register_buffer( action_bias, torch.tensor((env.action_space.high env.action_space.low) / 2.0, dtypetorch.float32) ) # 前向时tanh 输出 [-1,1]再映射回真实区间 return x * self.action_scale self.action_bias这一点非常重要MuJoCo 中并非所有环境动作区间都是 $[-1,1]$。例如Humanoid-v2为 $[-0.4, 0.4]$、InvertedPendulum-v2为 $[-3, 3]$、Pusher-v2为 $[-2, 2]$。action_bias/action_scale机制确保这类环境同样能被正确训练。探索噪声同样以action_bias为中心、以action_scale * exploration_noise为尺度采样。6.1 训练主循环的工程要点在 cleanrl/ddpg_continuous_action.py 的训练循环中还有两个容易遗漏但关键的细节final_observation处理当 episode 因截断truncation结束时用infos[final_observation]覆盖next_obs后再写入回放缓冲区避免将截断状态误当作终止状态参与 Bellman 更新延迟更新Actor 与目标网络每policy_frequency默认 2步才更新一次目标网络采用软更新target_param tau * param (1 - tau) * target_param这借鉴了 TD3 的思想以稳定训练。七、JAX 版本更快的后端与工程实现要点cleanrl/ddpg_continuous_action_jax.py 用 JAX、Flax、Optax 重写了同一算法文档与脚本中说明其在相近硬件下约比 PyTorch 版快 2.5~4 倍若关闭--capture-video开销则加速更明显。从源码看其工程实现有三个特点TrainState 携带目标参数自定义TrainState增加target_params字段与在线参数一起被优化器状态管理jax.jit编译训练函数Critic 更新update_critic与 Actor 更新update_actor均用jax.jit编译其中 Actor 更新通过optax.incremental_update(params, target_params, tau)完成目标网络的软更新模型序列化保存模型时使用flax.serialization.to_bytes将 Actor/Critic 参数序列化写入.cleanrl_model文件。JAX 版与 PyTorch 版共享相同的超参默认值3e-4学习率、256batch、0.005tau 等因此两份实现可以直接对照评测——这正是上一节评测表成立的前提。八、训练结束后的评估与模型导出两个实现都支持--save-model与--upload-model--save-model训练结束后将(actor.state_dict(), qf1.state_dict())PyTorch 版保存为runs/{run_name}/{exp_name}.cleanrl_model随后调用 cleanrl_utils/evals/ddpg_eval.py 中的evaluate()进行 10 个 episode 的确定性评测叠加exploration_noise后再裁剪并把eval/episodic_return写入 TensorBoard。评测时加载模型后调用actor.eval()仅用 Actor 前向推理--upload-model通过 cleanrl_utils/huggingface.py 的push_to_hub()自动创建 Hugging Face 仓库命名规则{env_id}-{exp_name}-seed{seed}上传模型权重、自动生成的模型卡与评测视频便于后续用python -m cleanrl_utils.enjoy --exp-name ddpg_continuous_action --env-id Hopper-v4直接加载玩耍。九、绘制学习曲线benchmark/ddpg_plot.sh 解读评测表中的平均回报来自学习曲线绘制脚本 benchmark/ddpg_plot.sh 使用openrlbenchmark.rlops拉取 wandb 上 tag 为pr-424的 run 并生成对比图python -m openrlbenchmark.rlops \ --filters ?weopenrlbenchmarkwpncleanrlceikenv_idcenexp_namemetriccharts/episodic_return \ ddpg_continuous_action?tagpr-424 \ ddpg_continuous_action_jax?tagpr-424 \ --env-ids HalfCheetah-v4 Walker2d-v4 Hopper-v4 InvertedPendulum-v4 Humanoid-v4 Pusher-v4 \ --no-check-empty-runs \ --pc.ncols 3 \ --pc.ncols-legend 2 \ --output-filename benchmark/cleanrl/ddpg_jax \ --scan-history其中--filters指定数据来源wandb 实体openrlbenchmark、项目cleanrl、指标charts/episodic_return?tagpr-424锁定对应版本的 run--scan-history用于重新扫描历史。这也意味着只要你自己用--track跑过基准就可以用同样的命令绘制自己的对比曲线。十、结果解读时的注意事项评测文档对结果差异给出了两点明确提醒解读数据时务必留意环境版本差异CleanRL 使用 Gymnasium MuJoCov4环境而参考实现Fujimoto et al. 2018报告的是早已废弃的 gym MuJoCov1环境两者动力学细节不同不能直接横向比较绝对值评测口径差异参考实现采用确定性评测无探索噪声训练结束后单独评测而 CleanRL 报告的是训练过程中的回合回报且策略在每个环境步之间持续更新因此 Walker2d、Hopper 等环境上的数字略低于参考实现属于预期现象而非实现缺陷。十一、快速冒烟验证测试用例仓库的 tests/test_mujoco.py 提供了极小的冒烟测试可以在秒级验证两份实现与模型导出链路均正常# 验证 PyTorch / JAX 训练主循环105 步 python cleanrl/ddpg_continuous_action.py --env-id Hopper-v4 --learning-starts 100 --batch-size 32 --total-timesteps 105 python cleanrl/ddpg_continuous_action_jax.py --env-id Hopper-v4 --learning-starts 100 --batch-size 32 --total-timesteps 105 # 验证 --save-model 与评测链路 python cleanrl/ddpg_continuous_action.py --save-model --env-id Hopper-v4 --learning-starts 100 --batch-size 32 --total-timesteps 105 python cleanrl/ddpg_continuous_action_jax.py --save-model --env-id Hopper-v4 --learning-starts 100 --batch-size 32 --total-timesteps 105这组命令同样适合在动手跑完整基准前先确认本机环境MuJoCo、PyTorch/JAX、tyro 参数解析一切就绪。十二、总结围绕 docs/benchmark/ddpg.md 这份评测数据本文还原了它的完整生产链路两个单文件实现PyTorch 与 JAX/Flax、6 个 MuJoCo v4 环境 × 3 种子的批量基准脚本 benchmark/ddpg.sh、TensorBoard/wandb 指标体系、与原始论文的工程差异、以及评测与绘图脚本。无论你是想复现这份数据、在自有环境上对比 DDPG 与 TD3/SAC 的表现还是评估 JAX 后端的加速收益都可以直接以仓库中的命令为起点把这张表变成你自己的实验。延伸阅读算法全景文档docs/rl-algorithms/ddpg.md含两个实现的完整用法、指标公式与实验曲线PyTorch 实现cleanrl/ddpg_continuous_action.pyJAX 实现cleanrl/ddpg_continuous_action_jax.py经验回放缓冲区实现cleanrl_utils/buffers.py来自 stable-baselines3handle_timeout_termination等细节可在此查阅评测脚本cleanrl_utils/evals/ddpg_eval.py 与 cleanrl_utils/evals/ddpg_jax_eval.py基准与绘图脚本benchmark/ddpg.sh、benchmark/ddpg_plot.sh相关算法对照实现cleanrl/td3_continuous_action.py、cleanrl/sac_continuous_action.py【免费下载链接】cleanrlHigh-quality single file implementation of Deep Reinforcement Learning algorithms with research-friendly features (PPO, DQN, C51, DDPG, TD3, SAC, PPG)项目地址: https://gitcode.com/GitHub_Trending/cl/cleanrl创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考