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 的表现:
- Clipped Double Q-Learning(截断的双 Q 学习):同时学习两个 Q 网络
qf1与qf2,在计算目标值时取两者的最小值,抑制 Q 值的过估计(overestimation); - Delayed Policy Updates(延迟的策略更新):critic(Q 网络)每个时间步都更新,而 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/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-id | Hopper-v4 | Gymnasium MuJoCo 环境 ID |
--total-timesteps | 1000000 | 总训练时间步数 |
--learning-rate | 3e-4 | Actor 与 Critic 优化器的学习率 |
--buffer-size | 1000000 | 经验回放缓冲区容量 |
--gamma | 0.99 | 折扣因子 |
--tau | 0.005 | 目标网络软更新系数(target smoothing coefficient) |
--batch-size | 256 | 每次从回放缓冲区采样的批大小 |
--policy-noise | 0.2 | 目标策略平滑正则化的噪声尺度 |
--exploration-noise | 0.1 | 训练时叠加在动作上的探索高斯噪声尺度 |
--learning-starts | 25000 | 开始学习的预热时间步数(此前只做随机探索) |
--policy-frequency | 2 | 策略(actor)更新频率:每 N 步更新一次(延迟更新) |
--noise-clip | 0.5 | 目标策略噪声的裁剪范围[-0.5, 0.5] |
--capture-video | False | 是否录制训练视频到videos/目录 |
--save-model/--upload-model | False | 是否保存模型 / 上传模型到 Hugging Face Hub |
--track | False | 是否用 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_noise(policy_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_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_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, 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) / 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时的 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为基准改写,除上一节讨论的动作重缩放外,还有以下实现差异:
- 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.py | TD3.py(Fujimoto et al., 2018, Table 1) |
|---|---|---|
| HalfCheetah-v4 | 9583.22 ± 126.09 | 9636.95 ± 859.065 |
| Walker2d-v4 | 4057.59 ± 658.78 | 4682.82 ± 539.64 |
| Hopper-v4 | 3134.61 ± 360.18 | 3564.07 ± 114.74 |
| InvertedPendulum-v4 | 968.99 ± 25.80 | 1000.00 ± 0.00 |
| Humanoid-v4 | 5035.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_action与td3_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_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 / Actor),Actor 通过action_scale、action_bias字段完成与 PyTorch 版一致的动作重缩放; - 训练状态:自定义
TrainState在 Flax 训练状态基础上扩展了target_params字段(TrainState),用optax.incremental_update实现目标网络的软更新(update_actor); - JIT 编译:
actor.apply、qf.apply以及update_critic、update_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-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与达到同等回报所需时间,即可直观感受 2~4 倍的吞吐差异; - 环境适配:如需在
[-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 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),仅供参考