免费获取学习方案
ARTICLE DETAIL

资讯详情

深耕编程基础知识与建站技术分享的一线实战洞察。

CleanRL 中 TD3 算法的单文件实现与实战指南:从 Clipped Double Q-Learning 到 JAX 加速

CleanRL 中 TD3 算法的单文件实现与实战指南:从 Clipped Double Q-Learning 到 JAX 加速 CleanRL 中 TD3 算法的单文件实现与实战指南从 Clipped Double Q-Learning 到 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 仓库中的 td3.md 为核心骨架系统讲解 Twin Delayed Deep Deterministic Policy GradientTD3算法在连续控制任务上的两种单文件实现 ——td3_continuous_action.pyPyTorch 版与td3_continuous_action_jax.pyJAX 版。文章将覆盖 TD3 的三大改进技巧、完整命令行用法、可复现的超参数、日志指标解读、与参考实现的源码级差异分析以及基于 benchmark/td3.sh 的基准测试流程帮助读者快速上手并深入理解 TD3 的工程实现细节。1. TD3 算法概述TD3Twin Delayed Deep Deterministic Policy Gradient是深度强化学习DRL中面向连续控制任务的主流算法。它是对 DDPG 的扩展通过引入三种关键技术来显著缓解 actor-critic 方法中常见的函数逼近误差function approximation error问题从而在多数连续控制基准上取得明显优于 DDPG 的表现Clipped Double Q-Learning截断的双 Q 学习同时学习两个 Q 网络qf1与qf2在计算目标值时取两者的最小值抑制 Q 值的过估计overestimationDelayed Policy Updates延迟的策略更新criticQ 网络每个时间步都更新而 actor策略的更新频率更低默认每 2 步更新 1 次让 Q 函数先充分收敛再指导策略更新Target Policy Smoothing Regularization目标策略平滑正则化为下一状态的目标动作注入带裁剪的高斯噪声使 Q 函数对相似动作的输出更平滑降低对动作空间的过拟合。原始论文与参考资源Fujimoto, S., van Hoof, H., Meger, D. (2018).Addressing Function Approximation Error in Actor-Critic Methods.OpenAI Spinning Up 的Twin Delayed DDPG章节参考实现sfujim/TD31.1 实现变体总览CleanRL 在 cleanrl/ 目录下提供了两种单文件 TD3 实现变体实现描述cleanrl/td3_continuous_action.py基于 PyTorch 的实现适用于连续动作空间cleanrl/td3_continuous_action_jax.py基于 JAX / Flax / Optax 的实现适用于连续动作空间同硬件下约比 PyTorch 版快 2.54 倍两者遵循 CleanRL 一贯的单文件、零抽象、可直接运行的设计理念每个文件都可以独立完成从环境初始化、训练循环、日志记录到模型保存/上传的完整流程。2.td3_continuous_action.pyPyTorch 单文件实现td3_continuous_action.py 是 TD3 的 PyTorch 参考实现具有以下特性面向连续动作空间设计支持Box类型的低维特征观测空间observation space支持Box连续动作空间通过assert isinstance(envs.single_action_space, gym.spaces.Box)在启动时校验动作空间类型非连续动作空间会直接报错见 td3_continuous_action.py。2.1 安装与运行使用包管理器安装 MuJoCo 相关依赖后即可运行pyproject.toml中已声明mujoco可选依赖组 uvpoetry 风格bash uv pip install .[mujoco] uv run python cleanrl/td3_continuous_action.py --help uv run python cleanrl/td3_continuous_action.py --env-id Hopper-v4 pipbash pip install -r requirements/requirements-mujoco.txt python cleanrl/td3_continuous_action.py --help python cleanrl/td3_continuous_action.py --env-id Hopper-v4 其中--help会通过tyro.cli(Args)自动生成完整的参数说明Args数据类定义在 td3_continuous_action.py。首次运行会生成runs/{env_id}__{exp_name}__{seed}__{timestamp}目录TensorBoard 日志、训练视频开启--capture-video时与模型文件开启--save-model时都会保存其中。2.2 核心超参数详解以下参数直接决定 TD3 的训练行为均通过命令行--参数名 值覆盖默认值参数默认值说明--env-idHopper-v4Gymnasium MuJoCo 环境 ID--total-timesteps1000000总训练时间步数--learning-rate3e-4Actor 与 Critic 优化器的学习率--buffer-size1000000经验回放缓冲区容量--gamma0.99折扣因子--tau0.005目标网络软更新系数target smoothing coefficient--batch-size256每次从回放缓冲区采样的批大小--policy-noise0.2目标策略平滑正则化的噪声尺度--exploration-noise0.1训练时叠加在动作上的探索高斯噪声尺度--learning-starts25000开始学习的预热时间步数此前只做随机探索--policy-frequency2策略actor更新频率每 N 步更新一次延迟更新--noise-clip0.5目标策略噪声的裁剪范围[-0.5, 0.5]--capture-videoFalse是否录制训练视频到videos/目录--save-model/--upload-modelFalse是否保存模型 / 上传模型到 Hugging Face Hub--trackFalse是否用 Weights Biases 跟踪实验代码中exp_name、seed、torch_deterministic、cuda、wandb_project_name、wandb_entity、hf_entity等实验管理参数与算法参数一并由 Args 数据类 声明完整参数列表可直接用--help查看。2.3 训练主循环中的关键机制从 训练主循环 可以看到 TD3 在 CleanRL 中的落地方式探索阶段global_step learning_starts时从动作空间均匀采样随机动作之后用actor输出动作并叠加N(0, action_scale * exploration_noise)的高斯噪声最后clip到动作空间边界目标值计算目标动作由target_actor生成并叠加裁剪后的噪声clipped_noisepolicy_noise乘以target_actor.action_scale后裁剪到±noise_clip再对两个目标 Q 网络输出取min得到 Bellman 目标r γ·min(qf1_next, qf2_next)Critic 更新对qf1、qf2分别计算 MSE 损失并求和后反向传播更新即qf_loss qf1_loss qf2_loss延迟的策略更新仅当global_step % policy_frequency 0时才更新 actor并同步用τ做目标网络的软更新soft update终止处理handle_timeout_terminationFalse表示 Gymnasium 中由truncation触发的回合结束不按终止termination处理避免把“超时截断”误当成真正失败这一细节在 ReplayBuffer 初始化 处体现。2.4 网络结构与动作空间重缩放Actor 与 Critic 均为两层 256 隐藏单元的 MLP。与参考实现 sfujim/TD3 不同的是CleanRL 将两个 Q 网络拆成两个独立的QNetwork对象qf1/qf2参考实现则用一个Critic类同时包含两个 Q 网络二者的更新逻辑在数学上完全等价见 td3_continuous_action.py。更重要的是CleanRL 版本通过action_scale与action_bias两个 buffer 完成动作空间的重缩放见 Actor 定义class Actor(nn.Module): def __init__(self, env): ... # action rescaling self.register_buffer( action_scale, torch.tensor((env.single_action_space.high - env.single_action_space.low) / 2.0, dtypetorch.float32), ) self.register_buffer( action_bias, torch.tensor((env.single_action_space.high env.single_action_space.low) / 2.0, dtypetorch.float32), ) def forward(self, x): x F.relu(self.fc1(x)) x F.relu(self.fc2(x)) x torch.tanh(self.fc_mu(x)) return x * self.action_scale self.action_bias参考实现 sfujim/TD3 只支持以[-1, 1]为中心的对称动作空间其 Actor 输出为self.max_action * torch.tanh(self.l3(a))隐含假设上下界关于 0 对称。CleanRL 版则显式计算action_scale (high - low) / 2action_bias (high low) / 2从而把tanh输出从[-1, 1]线性映射到[low, high]对非对称或非[-1,1]边界的连续动作空间同样适用。同理训练与评估时叠加的探索噪声也以action_bias为中心、以action_scale * exploration_noise为尺度见 探索噪声采样 与 td3_eval.py 中的评估逻辑。为什么需要动作重缩放MuJoCo 环境并非所有动作空间都是[-1, 1]。例如Humanoid-v2的动作边界是[-0.4, 0.4]InvertedPendulum-v2是[-3.0, 3.0]Pusher-v2是[-2.0, 2.0]。如果像参考实现那样硬编码max_action假设这些环境上的策略输出会超出合法范围或无法覆盖完整动作区间。3. 日志指标解读运行python cleanrl/td3_continuous_action.py会自动将指标写入 TensorBoard以及开启--track时的 WB。各指标含义如下charts/episodic_return每个回合的累积回报episodic return用于评估策略的整体表现charts/episodic_length每个回合的步数charts/SPS每秒环境步数steps per second衡量训练吞吐量losses/qf1_loss当前 Q 值与目标 Q 值之间的均方误差MSE最小化单步时间差分误差见 qf1_loss 计算 [ J(\theta^Q) \mathbb{E}{(s,a,r,s) \sim \mathcal{D}}\left[\left(Q(s,a) - y\right)^2\right], \quad y r \gamma \min\left(Q{\theta_1}(s,a), Q_{\theta_2}(s,a)\right) ]losses/qf2_loss第二个 Q 网络的同一 MSE 损失与qf1_loss一起被优化losses/qf_loss(qf1_loss qf2_loss) / 2两个 Q 损失的平均值是查看整体 critic 收敛情况的便捷指标见 日志记录代码losses/actor_loss实现为-qf1(data.observations, actor(data.observations)).mean()即actor 基于观测计算的动作所对应的 Q 值的负平均。最小化该损失等价于沿确定性策略梯度更新 actor 参数Fujimoto et al., 2018, Algorithm 1 [ \nabla_{\phi} J(\phi) \left.N^{-1} \sum \nabla_a Q_{\theta_1}(s, a)\right|{a\pi\phi(s)} \nabla_{\phi} \pi_\phi(s) ]losses/qf1_values实现为qf1(data.observations, data.actions).view(-1)即回放缓冲区采样数据的平均 Q 值用于判断是否存在 Q 值的过估计/欠估计见 qf1_values 记录。注意actor_loss 只在global_step % policy_frequency 0的延迟更新步被计算与记录因此 TensorBoard 中losses/actor_loss的采样频率是 critic 损失的一半每 100 步记录一次时实际是每 200 步才有一次 actor 损失值。4. 与参考实现的差异分析CleanRL 的 td3_continuous_action.py 以 sfujim/TD3 的TD3.py为基准改写除上一节讨论的动作重缩放外还有以下实现差异Q 网络的组织方式CleanRL 使用两个独立对象qf1、qf2表示 Clipped Double Q-Learning 中的两个 Q 函数参考实现TD3.py则用一个Critic类同时包含两个 Q 网络二者数学上等价CleanRL 还额外维护qf1_target、qf2_target两个目标网络探索噪声的分布训练时叠加的高斯噪声N(0, action_scale * exploration_noise)以action_bias为中心、按动作空间尺度缩放参考实现以 0 为中心、以max_action缩放环境版本差异CleanRL 使用 Gymnasium 的 MuJoCov4环境如Hopper-v4参考实现使用已长期弃用的 gym MuJoCov1环境两者动力学实现存在差异这也是部分基准如 Walker2d数值不同的原因之一评估方式差异参考实现在main.py中用确定性评估不加探索噪声报告平均回合回报而 CleanRL 报告的是训练过程中、策略在环境步之间持续更新时的回合回报二者统计口径不同比较时需注意。4.1 环境动作空间参考gym MuJoCo v2/v4 常见值以下是文档中记录的常见 MuJoCo 环境观测与动作空间供参考验证动作重缩放的必要性Ant-v2 Observation space: Box(-inf, inf, (111,), float64) Action space: Box(-1.0, 1.0, (8,), float32) HalfCheetah-v2 Observation space: Box(-inf, inf, (17,), float64) Action space: Box(-1.0, 1.0, (6,), float32) Hopper-v2 Observation space: Box(-inf, inf, (11,), float64) Action space: Box(-1.0, 1.0, (3,), float32) Humanoid-v2 Observation space: Box(-inf, inf, (376,), float64) Action space: Box(-0.4, 0.4, (17,), float32) InvertedDoublePendulum-v2 Observation space: Box(-inf, inf, (11,), float64) Action space: Box(-1.0, 1.0, (1,), float32) InvertedPendulum-v2 Observation space: Box(-inf, inf, (4,), float64) Action space: Box(-3.0, 3.0, (1,), float32) Pusher-v2 Observation space: Box(-inf, inf, (23,), float64) Action space: Box(-2.0, 2.0, (7,), float32) Reacher-v2 Observation space: Box(-inf, inf, (11,), float64) Action space: Box(-1.0, 1.0, (2,), float32) Swimmer-v2 Observation space: Box(-inf, inf, (8,), float64) Action space: Box(-1.0, 1.0, (2,), float32) Walker2d-v2 Observation space: Box(-inf, inf, (17,), float64) Action space: Box(-1.0, 1.0, (6,), float32)5. 基准测试与实验结果5.1 运行官方基准仓库提供了完整的基准测试脚本 benchmark/td3.sh基于cleanrl_utils.benchmark批量跑 6 个 MuJoCo 环境、3 个随机种子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/td3_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_template该脚本默认基于 SLURM 集群提交任务每个任务 1 块 GPU、10 个 CPU、18 个并发 worker如果在本地单机运行可去掉--slurm-*参数并调低--workers。5.2 PyTorch 版的平均回合回报3 个随机种子为了验证实现质量文档将 CleanRL 结果与 Fujimoto et al. (2018) 表 1 的报告值对比环境td3_continuous_action.pyTD3.pyFujimoto et al., 2018, Table 1HalfCheetah-v49583.22 ± 126.099636.95 ± 859.065Walker2d-v44057.59 ± 658.784682.82 ± 539.64Hopper-v43134.61 ± 360.183564.07 ± 114.74InvertedPendulum-v4968.99 ± 25.801000.00 ± 0.00Humanoid-v45035.36 ± 21.67不可用Pusher-v4-30.92 ± 1.05不可用几点对比注意事项CleanRL 使用 Gymnasium MuJoCov4参考实现使用v1环境动力学不一致导致数值天然存在偏差Walker2d 上的差距可能源于 gym v1/v4 的动力学差异参考实现社区的 issue 亦有讨论且 v1 环境已长期弃用、难以精确复现评估口径不同参考实现用确定性评估报告结果CleanRL 报告训练过程中的回合回报且策略在环境步之间持续更新。5.3 绘制学习曲线使用 benchmark/td3_plot.sh 中基于openrlbenchmark.rlops的绘图命令可以从 WB 拉取td3_continuous_action与td3_continuous_action_jax两个实验组tagpr-424的charts/episodic_return指标并生成对比图python -m openrlbenchmark.rlops \ --filters ?weopenrlbenchmarkwpncleanrlceikenv_idcenexp_namemetriccharts/episodic_return \ td3_continuous_action?tagpr-424 \ td3_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/td3 \ --scan-history仓库的 docs/rl-algorithms/td3-jax/ 目录保留了 JAX 版在 HalfCheetah、Walker2d、Hopper 三个环境上的学习曲线图以训练步数/训练时间为横轴可用于观察两种实现的收敛行为与速度差异。5.4 模型保存与评估训练结束后可用--save-model保存模型runs/{run_name}/{exp_name}.cleanrl_model并自动调用 cleanrl_utils/evals/td3_eval.py 中的evaluate函数进行 10 个回合的确定性评估评估时同样叠加exploration_noise尺度的高斯噪声并裁剪到动作空间边界评估回报写入eval/episodic_return。--upload-model则把模型与评估视频推送到 Hugging Face Hub相关代码。6.td3_continuous_action_jax.pyJAX 加速版td3_continuous_action_jax.py 是 TD3 的 JAX 移植版使用 JAX、Flax、Optax 替代 PyTorch在相同硬件下训练吞吐量约为 PyTorch 版的 2.54 倍如果关闭--capture-video的录制开销加速比更高。其余特性与 PyTorch 版一致支持连续动作空间、Box观测空间、Box动作空间并同样实现了动作重缩放。6.1 安装与运行 uvpoetry 风格bash uv pip install .[mujoco, jax] uv run python cleanrl/td3_continuous_action_jax.py --help uv run python cleanrl/td3_continuous_action_jax.py --env-id Hopper-v4 pipbash pip install -r requirements/requirements-mujoco.txt pip install -r requirements/requirements-jax.txt python cleanrl/td3_continuous_action_jax.py --help python cleanrl/td3_continuous_action_jax.py --env-id Hopper-v4 JAX 版的核心参数与 PyTorch 版完全一致Args 数据类 中env_id、total_timesteps、learning_rate、buffer_size、gamma、tau、batch_size、policy_noise、exploration_noise、learning_starts、policy_frequency、noise_clip等默认值均相同可直接迁移已有超参数配置。6.2 JAX 版实现要点从 td3_continuous_action_jax.py 的源码可以看到 JAX 移植的主要设计模型定义QNetwork与Actor用flax.linen定义QNetwork / ActorActor 通过action_scale、action_bias字段完成与 PyTorch 版一致的动作重缩放训练状态自定义TrainState在 Flax 训练状态基础上扩展了target_params字段TrainState用optax.incremental_update实现目标网络的软更新update_actorJIT 编译actor.apply、qf.apply以及update_critic、update_actor均用jax.jit编译把整个训练步采样、目标计算、梯度更新融合为加速执行图JIT 装饰随机数管理用jax.random.split显式管理 PRNG key 流为 critic 更新中的噪声生成拆分独立 keyupdate_critic探索噪声JAX 版以max_action * exploration_noise为高斯噪声尺度max_action envs.single_action_space.high[0]同样裁剪到动作空间边界探索动作生成模型保存使用flax.serialization.to_bytes序列化 actor 与两个 Q 网络的参数模型保存代码评估则由 cleanrl_utils/evals/td3_jax_eval.py 的evaluate完成其__main__演示了从 Hugging Face Hub 下载官方模型并评估的用法。6.3 实验与 TPU 历史结果JAX 版的基准命令同样位于 benchmark/td3.sh第 1219 行区别在于安装jax[cuda11_cudnn82]0.4.8并指定--command uv run python cleanrl/td3_continuous_action_jax.py --track。文档中的历史实验曾在 TPU 上运行结果与 GPU 上非常接近但运行时长因硬件而异——不同硬件之间不做严格横向对比这是刻意为之1) 在同一硬件上重跑全部实验计算成本过高2) 要求所有贡献者使用相同硬件既不现实也不具备包容性。大致预期是同硬件下 JAX 版有 24 倍的速度提升关闭--capture-video后加速比会更高。6.4 JAX 版学习曲线TPU 历史实验以下曲线展示了 JAX 版在 HalfCheetah-v2 上的训练收敛过程TPU 上的历史实验横轴分别为训练步数与训练时间纵轴为回合回报三条曲线对应不同硬件配置下的 CleanRL 实现最终均收敛到相近的回报水平验证了 JAX 移植的正确性7. 复现与验证建议快速冒烟测试仓库的 tests/test_mujoco.py 提供了 TD3 的冒烟测试命令通过缩短learning-starts、batch-size与total-timesteps快速验证代码可运行python cleanrl/td3_continuous_action.py --env-id Hopper-v4 --learning-starts 100 --batch-size 32 --total-timesteps 105 python cleanrl/td3_continuous_action_jax.py --env-id Hopper-v4 --learning-starts 100 --batch-size 32 --total-timesteps 105正式训练用默认超参数运行 100 万步约数小时取决于硬件观察charts/episodic_return是否单调上升并收敛监控 Q 值偏差训练中定期查看losses/qf1_values与losses/qf2_values——若远高于实际回报说明存在过估计若两者出现显著分歧说明双 Q 机制未正常工作对比两种实现同硬件下分别运行 PyTorch 版与 JAX 版对比charts/SPS与达到同等回报所需时间即可直观感受 24 倍的吞吐差异环境适配如需在[-1,1]之外或非对称动作空间的环境如Humanoid-v4、InvertedPendulum-v4、Pusher-v4上训练无需修改代码——action_scale/action_bias会自动适配动作边界。参考资料Fujimoto, S., van Hoof, H., Meger, D. (2018).Addressing Function Approximation Error in Actor-Critic Methods.ArXiv, abs/1802.09477.OpenAI Spinning Up in Deep RLTwin Delayed DDPG。参考实现sfujim/TD3CleanRL 的 td3_continuous_action.py 即基于其TD3.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),仅供参考
返回列表