CleanRL 中 TD3 算法的单文件实现与实战指南:从 Clipped Double Q-Learning 到 JAX 加速
2026/9/15 15:27:48 网站建设 项目流程

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 Gradient(TD3)算法在连续控制任务上的两种单文件实现 ——td3_continuous_action.py(PyTorch 版)与td3_continuous_action_jax.py(JAX 版)。文章将覆盖 TD3 的三大改进技巧、完整命令行用法、可复现的超参数、日志指标解读、与参考实现的源码级差异分析,以及基于 benchmark/td3.sh 的基准测试流程,帮助读者快速上手并深入理解 TD3 的工程实现细节。

1. TD3 算法概述

TD3(Twin Delayed Deep Deterministic Policy Gradient)是深度强化学习(DRL)中面向连续控制任务的主流算法。它是对 DDPG 的扩展,通过引入三种关键技术来显著缓解 actor-critic 方法中常见的函数逼近误差(function approximation error)问题,从而在多数连续控制基准上取得明显优于 DDPG 的表现:

  1. Clipped Double Q-Learning(截断的双 Q 学习):同时学习两个 Q 网络qf1qf2,在计算目标值时取两者的最小值,抑制 Q 值的过估计(overestimation);
  2. Delayed Policy Updates(延迟的策略更新):critic(Q 网络)每个时间步都更新,而 actor(策略)的更新频率更低(默认每 2 步更新 1 次),让 Q 函数先充分收敛再指导策略更新;
  3. 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/TD3

1.1 实现变体总览

CleanRL 在 cleanrl/ 目录下提供了两种单文件 TD3 实现:

变体实现描述
cleanrl/td3_continuous_action.py基于 PyTorch 的实现,适用于连续动作空间
cleanrl/td3_continuous_action_jax.py基于 JAX / Flax / Optax 的实现,适用于连续动作空间,同硬件下约比 PyTorch 版快 2.5~4 倍

两者遵循 CleanRL 一贯的单文件、零抽象、可直接运行的设计理念,每个文件都可以独立完成从环境初始化、训练循环、日志记录到模型保存/上传的完整流程。

2.td3_continuous_action.py:PyTorch 单文件实现

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可选依赖组):

=== "uv(poetry 风格)"

```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 ```

=== "pip"

```bash 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_nameseedtorch_deterministiccudawandb_project_namewandb_entityhf_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 更新:对qf1qf2分别计算 MSE 损失并求和后反向传播更新,即qf_loss = qf1_loss + qf2_loss
  • 延迟的策略更新:仅当global_step % policy_frequency == 0时才更新 actor,并同步用τ做目标网络的软更新(soft update);
  • 终止处理handle_timeout_termination=False表示 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_scaleaction_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, dtype=torch.float32), ) self.register_buffer( "action_bias", torch.tensor((env.single_action_space.high + env.single_action_space.low) / 2.0, dtype=torch.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) / 2
  • action_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时的 W&B)。各指标含义如下:

  • 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为基准改写,除上一节讨论的动作重缩放外,还有以下实现差异:

  1. Q 网络的组织方式:CleanRL 使用两个独立对象qf1qf2表示 Clipped Double Q-Learning 中的两个 Q 函数,参考实现TD3.py则用一个Critic类同时包含两个 Q 网络,二者数学上等价;CleanRL 还额外维护qf1_targetqf2_target两个目标网络;
  2. 探索噪声的分布:训练时叠加的高斯噪声N(0, action_scale * exploration_noise)action_bias为中心、按动作空间尺度缩放,参考实现以 0 为中心、以max_action缩放;
  3. 环境版本差异:CleanRL 使用 Gymnasium 的 MuJoCov4环境(如Hopper-v4),参考实现使用已长期弃用的 gym MuJoCov1环境,两者动力学实现存在差异,这也是部分基准(如 Walker2d)数值不同的原因之一;
  4. 评估方式差异:参考实现在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.py(Fujimoto et al., 2018, Table 1)
HalfCheetah-v49583.22 ± 126.099636.95 ± 859.065
Walker2d-v44057.59 ± 658.784682.82 ± 539.64
Hopper-v43134.61 ± 360.183564.07 ± 114.74
InvertedPendulum-v4968.99 ± 25.801000.00 ± 0.00
Humanoid-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的绘图命令,可以从 W&B 拉取td3_continuous_actiontd3_continuous_action_jax两个实验组(tagpr-424)的charts/episodic_return指标,并生成对比图:

python -m openrlbenchmark.rlops \ --filters '?we=openrlbenchmark&wpn=cleanrl&ceik=env_id&cen=exp_name&metric=charts/episodic_return' \ 'td3_continuous_action?tag=pr-424' \ 'td3_continuous_action_jax?tag=pr-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.py:JAX 加速版

td3_continuous_action_jax.py 是 TD3 的 JAX 移植版,使用 JAX、Flax、Optax 替代 PyTorch,在相同硬件下训练吞吐量约为 PyTorch 版的 2.5~4 倍(如果关闭--capture-video的录制开销,加速比更高)。其余特性与 PyTorch 版一致:支持连续动作空间、Box观测空间、Box动作空间,并同样实现了动作重缩放。

6.1 安装与运行

=== "uv(poetry 风格)"

```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 ```

=== "pip"

```bash 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_idtotal_timestepslearning_ratebuffer_sizegammataubatch_sizepolicy_noiseexploration_noiselearning_startspolicy_frequencynoise_clip等默认值均相同),可直接迁移已有超参数配置。

6.2 JAX 版实现要点

从 td3_continuous_action_jax.py 的源码可以看到 JAX 移植的主要设计:

  • 模型定义QNetworkActorflax.linen定义(QNetwork / Actor),Actor 通过action_scaleaction_bias字段完成与 PyTorch 版一致的动作重缩放;
  • 训练状态:自定义TrainState在 Flax 训练状态基础上扩展了target_params字段(TrainState),用optax.incremental_update实现目标网络的软更新(update_actor);
  • JIT 编译actor.applyqf.apply以及update_criticupdate_actor均用jax.jit编译,把整个训练步(采样、目标计算、梯度更新)融合为加速执行图(JIT 装饰);
  • 随机数管理:用jax.random.split显式管理 PRNG key 流,为 critic 更新中的噪声生成拆分独立 key(update_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(第 12~19 行),区别在于安装jax[cuda11_cudnn82]==0.4.8并指定--command "uv run python cleanrl/td3_continuous_action_jax.py --track"。文档中的历史实验曾在 TPU 上运行,结果与 GPU 上非常接近,但运行时长因硬件而异——不同硬件之间不做严格横向对比,这是刻意为之:1) 在同一硬件上重跑全部实验计算成本过高;2) 要求所有贡献者使用相同硬件既不现实也不具备包容性。大致预期是同硬件下 JAX 版有 2~4 倍的速度提升,关闭--capture-video后加速比会更高。

6.4 JAX 版学习曲线(TPU 历史实验)

以下曲线展示了 JAX 版在 HalfCheetah-v2 上的训练收敛过程(TPU 上的历史实验,横轴分别为训练步数与训练时间,纵轴为回合回报),三条曲线对应不同硬件配置下的 CleanRL 实现,最终均收敛到相近的回报水平,验证了 JAX 移植的正确性:

7. 复现与验证建议

  • 快速冒烟测试:仓库的 tests/test_mujoco.py 提供了 TD3 的冒烟测试命令,通过缩短learning-startsbatch-sizetotal-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_valueslosses/qf2_values——若远高于实际回报,说明存在过估计;若两者出现显著分歧,说明双 Q 机制未正常工作;
  • 对比两种实现:同硬件下分别运行 PyTorch 版与 JAX 版,对比charts/SPS与达到同等回报所需时间,即可直观感受 2~4 倍的吞吐差异;
  • 环境适配:如需在[-1,1]之外或非对称动作空间的环境(如Humanoid-v4InvertedPendulum-v4Pusher-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 RL:Twin Delayed DDPG
  • 参考实现:sfujim/TD3(CleanRL 的 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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询