CleanRL 与 EnvPool 基准测试:PPO 在 Atari-v5 上的 PyTorch 与 JAX 三种实现对比与复现

发布时间:2026/9/15 18:53:09
CleanRL 与 EnvPool 基准测试:PPO 在 Atari-v5 上的 PyTorch 与 JAX 三种实现对比与复现 CleanRL 与 EnvPool 基准测试PPO 在 Atari-v5 上的 PyTorch 与 JAX 三种实现对比与复现【免费下载链接】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 仓库中 EnvPool 基准测试docs/benchmark/ppo_envpool.md为核心围绕 PPO 算法在 Atari-v5 环境上的三种实现——PyTorch 版ppo_atari_envpool.py、JAX 版ppo_atari_envpool_xla_jax.py及其scan优化版ppo_atari_envpool_xla_jax_scan.py——展开。读者将了解到三份基准结果表的完整数据与解读方法、各实现的源码级差异与性能瓶颈所在以及如何通过cleanrl_utils.benchmark与 benchmark/ppo.sh 一键复现这些实验并验证模型正确性。一、为什么用 EnvPool 跑 Atari-v5环境与接口背景传统 Atari 实验如PongNoFrameskip-v4基于 Gym 与单进程模拟器采样吞吐有限。CleanRL 的 EnvPool 系列使用 EnvPool 提供的v5 环境如Pong-v5、BeamRider-v5、Breakout-v5它具有两个关键特性向量化并行envpool.make(env_id, env_typegym, num_envsargs.num_envs, ...)一次性创建多个并行环境通过统一的向量接口同时推进配合episodic_lifeTrue失一条命即截断用于训练与reward_clipTrue奖励裁剪到 [-1, 1]等经典 Atari 预处理见 cleanrl/ppo_atari_envpool.py 与 cleanrl/ppo_atari_envpool_xla_jax.py。XLA 原生接口JAX 版通过envs.xla()拿到handle, recv, send, step_env将环境推进直接嵌入 JIT 编译的rollout函数实现“采样-训练”全链路在 GPU 上零 Python 开销执行这是吞吐量远超 PyTorch 版的根本原因。在依赖层面仓库通过uv pip install .[envpool]安装 EnvPool锁定的版本为envpool0.6.6见 requirements/requirements-envpool.txtJAX 版本额外需要.[envpool, jax]。二、基准结果总表三种实现的三环境性能对比docs/benchmark/ppo_envpool.md记录了 tag 为pr-424的三次实验每项 3 个随机种子均值 ± 标准差对比对象为ppo_atari_envpool_xla_jax、ppo_atari_envpool_xla_jax_scan与ppo_atari_envpool。所有实验均在默认超参数见下节下训练 1000 万步环境openrlbenchmark/cleanrl/ppo_atari_envpool_xla_jax (pr-424)openrlbenchmark/cleanrl/ppo_atari_envpool_xla_jax_scan (pr-424)openrlbenchmark/cleanrl/ppo_atari_envpool (pr-424)Pong-v520.82 ± 0.2120.52 ± 0.3220.45 ± 0.09BeamRider-v52678.73 ± 426.422860.61 ± 801.302501.85 ± 210.52Breakout-v5420.92 ± 16.75423.90 ± 5.49211.24 ± 151.84从数据可以得出几个可验证的观察训练效果等价性在 Pong-v5 上三种实现都收敛到约 20 分接近理论满分说明不同框架后端不改变算法本身的行为JAX 版 Breakout-v5 的方差±16.75 / ±5.49显著小于 PyTorch 版±151.84可推断与XLA_PYTHON_CLIENT_MEM_FRACTION等确定性环境配置见下文源码分析有关但具体因果还需更多重复实验确认。方差主要来自环境与种子BeamRider-v5 上三种实现的均值都落在 25002900 区间但标准差高达 210800说明该环境的回报本身对种子敏感框架选择不是主要变量。三、运行时间对照JAX 提速的量化证据配套文件 docs/benchmark/ppo_envpool_runtimes.md 记录了相同 3 次实验的训练耗时表中数值为单次实验完成 1000 万步训练所需时间数值越小越快环境ppo_atari_envpool_xla_jax (pr-424)ppo_atari_envpool_xla_jax_scan (pr-424)ppo_atari_envpool (pr-424)Pong-v534.323734.701178.375BeamRider-v537.107637.2449182.944Breakout-v539.57639.775151.384结论非常清晰JAX 版本比 PyTorch EnvPool 版本快约 4.55 倍例如 Breakout-v5 从 151.384 降到 39.576而xla_jax与xla_jax_scan两种 JAX 写法在吞吐上几乎持平差异小于 1%。这意味着如果追求极致的训练吞吐JAX 后端是决定性因素scan改写主要带来代码结构上的收益见第五节而非额外的速度提升。四、EnvPool 与经典 Gym Atari 的横向对比为了评估 EnvPool v5 路线相对仓库传统ppo_atari.pyGym NoFrameskip-v4的取舍仓库还提供了两份对照表。4.1 训练效果对比docs/benchmark/ppo_atari_envpool.md环境openrlbenchmark/cleanrl/ppo_atari_envpool (pr-424)openrlbenchmark/cleanrl/ppo_atari (pr-424)Pong-v520.45 ± 0.0920.36 ± 0.20BeamRider-v52501.85 ± 210.521915.93 ± 484.58Breakout-v5211.24 ± 151.84414.66 ± 28.094.2 运行时间对比docs/benchmark/ppo_atari_envpool_runtimes.md环境ppo_atari_envpool (pr-424)ppo_atari (pr-424)Pong-v5178.375281.071BeamRider-v5182.944284.941Breakout-v5151.384264.077解读要点吞吐优势稳定在三个环境上EnvPool 版耗时都是经典版的一半左右验证了向量化环境对 Atari 采样的加速效果。效果互有胜负EnvPool 版在 BeamRider-v5 上均值更高2501.85 vs 1915.93而经典版在 Breakout-v5 上表现更稳定414.66 ± 28.09 vs 211.24 ± 151.84。由于v5与NoFrameskip-v4在环境实现、随机数与预处理细节上存在差异这一差别不能直接归因于算法质量只能说明两者是“等价但不同”的实验协议。五、三种实现的源码级拆解仓库在 cleanrl/ 目录下维护了三个独立文件实现同一套 PPO 超参数共用相同的Args数据类核心参数如下以 cleanrl/ppo_atari_envpool_xla_jax.py 为准三个文件完全一致参数默认值说明env_idBreakout-v5Atari v5 环境 idtotal_timesteps10,000,000总训练步数learning_rate2.5e-4Adam 初始学习率anneal_lrTrue时线性退火到 0num_envs8并行环境数num_steps128每次 rollout 每环境的步数batch_size num_envs * num_steps 1024gamma/gae_lambda0.99 / 0.95折扣因子与 GAE 参数num_minibatches/update_epochs4 / 4每批切 4 个 minibatch每批数据复用 4 个 epoch即每次迭代 16 次梯度更新clip_coef/clip_vloss0.1 / TruePPO 裁剪系数价值函数也使用裁剪损失ent_coef/vf_coef0.01 / 0.5熵正则系数与价值损失系数max_grad_norm0.5梯度全局范数裁剪norm_adv/target_klTrue / None优势归一化target_kl非空时启用早停5.1 PyTorch 版ppo_atari_envpool.py该版本结构最接近经典单文件 PPORecordEpisodeStatistics包装器逐帧累计回合回报并写入infos[r]、infos[l]cleanrl/ppo_atari_envpool.py策略网络是标准 CNN4→32×8×8/s4→64×4×4/s2→64×3×3/s1→512actor 输出层std0.01正交初始化、critic 输出层std1cleanrl/ppo_atari_envpool.py。训练循环中PPO 损失、GAE、minibatch 更新均在 Python 层用 PyTorch 张量完成环境交互则通过 EnvPool 向量接口envs.step(action.cpu().numpy())进行cleanrl/ppo_atari_envpool.py。5.2 JAX 版ppo_atari_envpool_xla_jax.py该版本把几乎整个训练循环搬进 XLA 图内文件头部设置了三个环境变量来保证 JAX 在 GPU 上的稳定性XLA_PYTHON_CLIENT_MEM_FRACTION0.6缓解显存 OOM、TF_XLA_FLAGS--xla_gpu_autotune_level2 --xla_gpu_deterministic_reductions与TF_CUDNN DETERMINISTIC1保证确定性归约见 cleanrl/ppo_atari_envpool_xla_jax.py。网络用 Flax 定义AgentParams以FrozenDict持有network/actor/critic三份参数cleanrl/ppo_atari_envpool_xla_jax.py优化器为optax.chain(clip_by_global_norm, inject_hyperparams(adam))学习率由linear_schedule按“累计梯度更新次数”线性退火cleanrl/ppo_atari_envpool_xla_jax.py。动作采样使用 Gumbel-softmax 技巧argmax(logits - log(-log(u)))GAE 通过纯函数compute_gae计算cleanrl/ppo_atari_envpool_xla_jax.pyminibatch 更新由jax.value_and_grad驱动rollout整体被jax.jit编译cleanrl/ppo_atari_envpool_xla_jax.py。5.3 scan 版ppo_atari_envpool_xla_jax_scan.pyxla_jax_scan版在 5.2 的基础上做了两处纯函数化改造以消除 Python 循环、减少内核启动次数GAE 用jax.lax.scan(compute_gae_once, ...)以reverseTrue倒序扫描实现cleanrl/ppo_atari_envpool_xla_jax_scan.py多 epoch、多 minibatch 的更新过程全部收敛进jax.lax.scan包括对 storage 的 shuffle、reshape 与逐 minibatch 的梯度应用cleanrl/ppo_atari_envpool_xla_jax_scan.py参考了 Brax 的实现模式额外支持--save-model序列化 Flax 参数到runs/{run_name}/*.cleanrl_model与--upload-model推送到 Hugging Face Hub并复用cleanrl_utils.evals.ppo_envpool_jax_eval做离线评测cleanrl/ppo_atari_envpool_xla_jax_scan.py。六、如何复现基准实验benchmark/ppo.sh是仓库的官方复现脚本其中与本文主题直接相关的三组命令如下。6.1 复现 PyTorch EnvPool 三环境基准uv pip install .[envpool] uv run python -m cleanrl_utils.benchmark \ --env-ids Pong-v5 BeamRider-v5 Breakout-v5 \ --command uv run python cleanrl/ppo_atari_envpool.py --track --capture_video \ --num-seeds 3 \ --workers 9 \ --slurm-gpus-per-task 1 \ --slurm-ntasks 1 \ --slurm-total-cpus 10 \ --slurm-template-path benchmark/cleanrl_1gpu.slurm_template对应 benchmark/ppo.sh。--num-seeds 3对应结果表中的 3 个种子--workers 9并行调度 9 个任务--track开启 WB 记录在--command内传给训练脚本--capture_video录制评测视频Slurm 相关参数按模板 benchmark/cleanrl_1gpu.slurm_template 提交到集群。6.2 复现 JAX 57 游戏全量基准uv pip install .[envpool, jax] uv run python -m cleanrl_utils.benchmark \ --env-ids Alien-v5 Amidar-v5 Assault-v5 Asterix-v5 Asteroids-v5 Atlantis-v5 ... Zaxxon-v5 \ --command uv run python ppo_atari_envpool_xla_jax.py --track --wandb-project-name envpool-atari --wandb-entity openrlbenchmark \ --num-seeds 3 \ --workers 9 \ --slurm-gpus-per-task 1 \ --slurm-ntasks 1 \ --slurm-total-cpus 10 \ --slurm-template-path benchmark/cleanrl_1gpu.slurm_template对应 benchmark/ppo.sh。--env-ids一次性列出 57 个 Atari-v5 游戏--wandb-project-name envpool-atari与--wandb-entity openrlbenchmark指定 WB 项目归属正是结果表中openrlbenchmark/envpool-atari/ppo_atari_envpool_xla_jax一列的来源。6.3 复现 scan 版uv pip install .[envpool, jax] python -m cleanrl_utils.benchmark \ --env-ids Pong-v5 BeamRider-v5 Breakout-v5 \ --command uv run python cleanrl/ppo_atari_envpool_xla_jax_scan.py --track --capture_video \ --num-seeds 3 \ --workers 9 \ --slurm-gpus-per-task 1 \ --slurm-ntasks 1 \ --slurm-total-cpus 10 \ --slurm-template-path benchmark/cleanrl_1gpu.slurm_template对应 benchmark/ppo.sh。6.4 单机直接运行如果不使用 Slurm 集群可以直接运行训练脚本本身例如复现 PyTorch 版 1000 万步训练uv run python cleanrl/ppo_atari_envpool.py --env-id Pong-v5 --total-timesteps 10000000 --track所有超参数均通过 tyro 从命令行传入cleanrl/ppo_atari_envpool.pybatch_size、minibatch_size、num_iterations在运行时由公式num_envs * num_steps、batch_size // num_minibatches、total_timesteps // batch_size自动推导。七、结果展示与可视化7.1 全量 57 游戏结果完整的 57 游戏基准结果位于 docs/benchmark/ppo_atari_envpool_xla_jax.md与openrlbenchmark/baselines/baselines-ppo2-cnnOpenAI Baselines PPO2逐游戏对比。其中 CleanRL 的 JAX 实现在一批游戏上表现突出例如 Assault-v56791.74 ± 420.03 vs 4878.67 ± 815.64、Atlantis-v53778458.33 ± 117680.68 vs 2036749.00 ± 95929.75、ChopperCommand-v55642.83 ± 802.34 vs 816.33 ± 114.14、DemonAttack-v529283.83 ± 7007.31 vs 13788.43 ± 1313.44、UpNDown-v5487495.41 ± 39751.49 vs 156143.70 ± 70620.88、YarsRevenge-v555757.68 ± 7467.49 vs 9394.97 ± 2743.74等在另一些游戏上则低于基线如 KungFuMaster-v5、NameThisGame-v5、VideoPinball-v5体现了 Atari 基准中常见的“单算法无绝对优势”现象。仓库为这套全量基准提供了可视化图表可直接用于论文或博客引用7.2 运行日志与指标训练过程中的charts/avg_episodic_return、charts/SPS、losses/policy_loss、losses/value_loss、losses/approx_kl、losses/clipfrac、losses/explained_variance等标量会被写入 TensorBoardruns/{run_name}并在--track时同步到 WBJAX 版额外记录charts/SPS_update每次迭代内吞吐与charts/avg_episodic_length。SPS每秒步数是衡量吞吐的核心指标JAX 版大幅领先的原因可直接从该指标读出。八、测试验证快速冒烟实验仓库在 tests/test_envpool.py 中为三种实现提供了完整的冒烟测试用小步数验证训练管线可端到端跑通# PyTorch EnvPool 版 python cleanrl/ppo_atari_envpool.py --num-envs 8 --num-steps 32 --total-timesteps 256 # JAX 版缩减 epoch/minibatch 以加快编译 python cleanrl/ppo_atari_envpool_xla_jax.py --num-envs 8 --num-steps 6 --update-epochs 1 --num-minibatches 1 --total-timesteps 256 # scan 版 模型保存/评测 python cleanrl/ppo_atari_envpool_xla_jax_scan.py --save-model --num-envs 8 --num-steps 6 --update-epochs 1 --num-minibatches 1 --total-timesteps 256对应 tests/test_envpool.py 中的test_ppo_atari_envpool、test_ppo_atari_envpool_xla_jax、test_ppo_atari_envpool_xla_jax_scan与test_ppo_atari_envpool_xla_jax_scan_eval。注意JAX 测试刻意使用--update-epochs 1 --num-minibatches 1是为了在保证训练逻辑完整执行的前提下显著缩短首次 JIT 编译与执行时间若要得到有统计意义的结果仍应回到第六节的 3 种子、1000 万步配置。九、总结与选型建议综合以上数据与源码可以给出如下可验证的选型参考追求吞吐与大规模消融优先选择ppo_atari_envpool_xla_jax(_scan)。JAX 版比 PyTorch EnvPool 快约 4.55 倍见第三节运行时间表且 57 游戏全量基准证明其训练效果与经典实现一致见 docs/benchmark/ppo_atari_envpool_xla_jax.md。追求可读性与教学ppo_atari_envpool.py保留了完整 Python 层训练循环结构最接近教科书 PPO便于逐行理解与二次开发。追求代码纯函数化xla_jax_scan用jax.lax.scan消除所有显式循环吞吐与xla_jax持平并额外支持模型保存与 Hugging Face 上传适合作为生产管线的起点。环境协议选择EnvPool v5 相比经典NoFrameskip-v4可将训练耗时减半见 docs/benchmark/ppo_atari_envpool_runtimes.md但两者在部分游戏上的回报存在差异发布结果时应明确标注所使用的环境协议。【免费下载链接】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),仅供参考

关于本文作者

来自尧图内容编辑团队

尧图内容编辑团队 内容团队

尧图内容编辑团队

本文由尧图网络内容编辑团队执笔。团队由资深项目经理、前端工程师与设计师组成,所有内容均来自亲手交付的真实项目,先讲清问题、再给出可落地的解法。尧图深耕北京网站建设十年,服务过京华建材集团、智造科技等各行业客户,把一线经验沉淀为可复用的行业观察。

  • 十年建站经验,覆盖建材、制造、服务、文创等
  • 项目经理把关选题与事实准确性
  • 工程师与设计师联合撰写专业细节
  • 统一编辑规范,保证文风与排版一致
  • 每月复盘转化数据,迭代选题方向

延伸阅读

相关资讯与近期热门内容

深度阅读推荐

建站决策前值得细读的三篇

网站改版的5个关键决策
2024-08-12

网站改版的5个关键决策

什么时候该改版、改到什么程度、如何避免流量掉光,京华建材集团改版复盘给出答案。

获取专属建站方案

看完文章,把您的行业与预算告诉我们,免费获取一份量身定制的官网建设方案与报价。

立即免费咨询