
CleanRL QDagger 实现解析基于教师策略与回放缓冲区复用加速 DQN 训练【免费下载链接】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/cleanrlQDaggerQ-learning with Distillation and AGgregation是 DQN 算法的一种扩展它复用先前计算好的结果——包括教师策略teacher policy与教师回放缓冲区teacher replay buffer——来辅助训练学生策略student policy从而避免从零开始学习显著提升样本效率并降低训练新策略的计算开销。本文以 CleanRL 仓库中 docs/rl-algorithms/qdagger.md 为骨架结合 cleanrl/qdagger_dqn_atari_impalacnn.py 与 cleanrl/qdagger_dqn_atari_jax_impalacnn.py 的完整源码系统讲解 QDagger 的核心原理、安装运行方式、三阶段训练流程、日志指标含义与实验结果并对比 PyTorch 与 JAX 两个实现变体的差异。QDagger 算法概述QDagger 源自 Agarwal、Schwarzer、Castro、Courville 与 Bellemare 在 2022 年发表的论文Reincarnating Reinforcement Learning: Reusing Prior Computation to Accelerate ProgressarXiv:2206.01626。其核心思想是与其每次训练新策略都从零开始fresh start不如转世reincarnate复用已有的计算成果让新策略以旧策略为教师在旧经验的基础上继续进化。在 CleanRL 的具体实现中QDagger 的工作方式为从 Hugging Face Hub 下载一个已经训练好的 DQN Atari 教师策略CleanRL 自己的dqn_atari/dqn_atari_jax模型用该教师策略在环境中采集经验构建教师回放缓冲区用教师回放缓冲区对学生策略进行离线offline训练损失函数同时包含标准 DQN 的 TD 损失与蒸馏distillation损失进入在线online训练阶段蒸馏系数根据学生与教师的回报差距自适应衰减最终学生策略完全接管环境交互。这样学生策略不仅从环境奖励中学习还从教师策略的动作分布中蒸馏知识相当于把教师已经掌握的技能迁移给新策略。已实现的两个变体CleanRL 针对 QDagger 提供了两个单文件实现均为 Atari 游戏场景使用来自 RainbowDQN 的 Impala-CNN 骨干网络与通用 Atari 预处理技术变体框架说明qdagger_dqn_atari_impalacnn.pyPyTorch基于 CleanRL 的dqn_atari教师策略qdagger_dqn_atari_jax_impalacnn.pyJAX / Flax / Optax基于 CleanRL 的dqn_atari_jax教师策略比 PyTorch 版本快约 25%–50%两个脚本都支持 Atari 的像素Box观测空间原始形状(210, 160, 3)与Discrete动作空间。从源码看二者结构高度对称共享相同的 Atari 预处理封装见 cleanrl_utils/atari_wrappers.py 中的NoopResetEnv、MaxAndSkipEnv、EpisodicLifeEnv、FireResetEnv、ClipRewardEnv再叠加ResizeObservation、GrayScaleObservation与FrameStack将画面缩放到 84×84 灰度并堆叠 4 帧以及相同的 ReplayBuffer改编自 stable-baselines3支持optimize_memory_usage内存优化模式。PyTorch 版本qdagger_dqn_atari_impalacnn.py功能特性适用于 Atari 游戏使用来自 RainbowDQN 的 Impala-CNN 骨干与通用 Atari 预处理教师策略使用 Hugging Facecleanrl仓库中的dqn_atari策略兼容 Atari 像素Box观测空间(210, 160, 3)与Discrete动作空间。安装与运行使用 uvpoetry 兼容方式安装并运行uv pip install .[atari] uv run python cleanrl/qdagger_dqn_atari_impalacnn.py --env-id BreakoutNoFrameskip-v4 uv run python cleanrl/qdagger_dqn_atari_impalacnn.py --env-id PongNoFrameskip-v4或使用 pip 安装 requirements/requirements-atari.txt 中的依赖pip install -r requirements/requirements-atari.txt python cleanrl/qdagger_dqn_atari_impalacnn.py --env-id BreakoutNoFrameskip-v4 python cleanrl/qdagger_dqn_atari_impalacnn.py --env-id PongNoFrameskip-v4运行后会通过 TensorBoard 自动记录各类指标回报、损失等日志输出在runs/{run_name}目录下若设置--track还会同步到 Weights and Biases。核心命令行参数除 CleanRL 单文件脚本通用的参数--seed、--capture-video、--save-model、--upload-model、--hf-entity等与标准 DQN 超参数外QDagger 专属参数定义在 Args 中如下表所示参数默认值说明teacher_policy_hf_repoNone教师策略的 Hugging Face 仓库为None时自动推导为cleanrl/{env_id}-dqn_atari-seed1teacher_model_exp_namedqn_atari教师模型的实验名对应仓库中的模型文件前缀teacher_eval_episodes10教师策略评估回合数teacher_steps500000用教师策略运行生成回放缓冲区的步数offline_steps500000学生策略在教师回放缓冲区上离线训练的步数temperature1.0QDagger 温度参数用于缩放 Q 值 logits标准 DQN 超参数同样重要例如env_id默认BreakoutNoFrameskip-v4、total_timesteps默认 10000000即 10M 步 / 40M 帧、learning_rate默认1e-4、gamma默认0.99、tau默认1.0硬更新、target_network_frequency默认 1000、batch_size默认 32、start_e/end_e默认1.0/0.01、exploration_fraction默认0.10、learning_starts默认 80000、train_frequency默认 4、buffer_size默认 1000000。三阶段训练流程源码级解析qdagger_dqn_atari_impalacnn.py的训练主循环清晰地分为三个阶段对应源码中的三段逻辑阶段一教师评估与数据收集。脚本通过hf_hub_download从teacher_policy_hf_repo下载教师权重加载到从 cleanrl/dqn_atari.py 导入的TeacherModel即dqn_atari.QNetwork中并置于eval模式随后用 cleanrl_utils/evals/dqn_eval.py 中的evaluate评估教师策略记录charts/teacher/avg_episodic_return。接着运行teacher_steps步让教师策略ε-greedyε 从 1.0 线性衰减到 0.01与环境交互并把(obs, real_next_obs, actions, rewards, terminations, infos)写入teacher_rbReplayBuffer容量默认 100 万。源码注释明确说明这里假设我们无法访问教师训练时留下的回放缓冲区参见原论文 Fig. A.19因此需要自行用教师策略重新采样生成。阶段二离线学生训练。在offline_steps步内每步从teacher_rb采样一个batch_size的小批量同时计算 TD 目标与蒸馏损失with torch.no_grad(): target_max, _ target_network(data.next_observations).max(dim1) td_target data.rewards.flatten() args.gamma * target_max * (1 - data.dones.flatten()) teacher_q_values teacher_model(data.observations) / args.temperature student_q_values q_network(data.observations) old_val student_q_values.gather(1, data.actions).squeeze() q_loss F.mse_loss(td_target, old_val) student_q_values student_q_values / args.temperature distill_loss torch.mean(kl_divergence_with_logits(teacher_q_values, student_q_values)) loss q_loss 1.0 * distill_loss注意离线阶段蒸馏系数固定为1.0。每 100000 步保存一次模型并评估学生策略记录charts/offline/avg_episodic_return。阶段三在线学生训练。学生策略接管环境交互total_timesteps步内继续向自己的回放缓冲区rb采样训练。此时蒸馏系数变成自适应if len(episodic_returns) 10: distill_coeff 1.0 else: distill_coeff max(1 - np.mean(episodic_returns) / np.mean(teacher_episodic_returns), 0) loss q_loss distill_coeff * distill_loss即学生策略回报越接近甚至超过教师蒸馏系数越小最终完全由 TD 损失主导实现从模仿教师到超越教师的平滑过渡。网络架构Impala-CNN两个脚本的网络结构均移植自 AIcrowd NeurIPS 2020 procgen starter-kit 的 Impala-CNN。PyTorch 版本定义了ResidualBlock两层 3×3 卷积 残差连接与ConvSequence3×3 卷积 3×3/stride 2 最大池化 两个残差块QNetwork依次堆叠输出通道为[16, 32, 32]的三个ConvSequence再接Flatten、ReLU、256 维全连接层与输出维度等于动作数的线性层前向时对输入先除以 255.0 归一化。JAX 版本用 Flaxnn.linen实现了结构完全一致的网络注意其输入先jnp.transpose调整为NHWC布局。日志指标说明运行脚本会自动记录以下指标对应原文档的全部指标项charts/episodic_return游戏的回合回报episodic returncharts/SPS每秒与环境交互的步数steps per secondlosses/td_loss时间步 $t$ 的 Q 值与贝尔曼更新目标之间的均方误差MSE即最小化单步时间差分。形式上$$ J(\theta^{Q}) \mathbb{E}_{(s,a,r,s) \sim \mathcal{D}} \big[ (Q(s, a) - y)^2 \big], $$其中贝尔曼更新目标为 $y r \gamma , Q^{}(s, a)$$\mathcal{D}$ 为回放缓冲区losses/q_values实现为qf1(data.observations, data.actions).view(-1)即回放缓冲区采样数据的平均 Q 值可用于判断是否存在高估或低估losses/distill_loss蒸馏损失即教师策略 $\pi_T$ 与学生策略 $\pi$ 之间的 KL 散度$$ L_{\text{distill}} \lambda_t \mathbb{E}_{(s,a,r,s) \sim \mathcal{D}} \left[ \sum_a \pi_T(a|s)\log\pi(a|s)\right] $$源码中由kl_divergence_with_logits见 qdagger_dqn_atari_impalacnn.py实现即-softmax(target_logits) * (log_softmax(prediction_logits) - log_softmax(target_logits))两者在除以temperature后计算在线阶段再对批量求均值charts/distill_coeff蒸馏损失系数 $\lambda_t$是教师策略 $\pi_T$ 与学生策略 $\pi$ 回报之比的函数$$ \lambda_t \mathbb{1}_{tt_0}\max(1 - G^{\pi}/G^{\pi_T}, 0) $$losses/loss总损失为 TD 损失与蒸馏损失之和$$ L_{\text{qdagger}} J(\theta^{Q}) L_{\text{distill}} $$charts/teacher/avg_episodic_return教师策略评估的平均回合回报charts/offline/avg_episodic_return离线训练阶段策略评估的平均回合回报。与原始论文的实现差异qdagger_dqn_atari_impalacnn.py基于 (Agarwal et al., 2022)但存在以下几点实现差异回放缓冲区来源原论文直接使用教师策略训练时保存的回放缓冲区数据而 CleanRL 的教师策略dqn_atari并不包含回放缓冲区数据因此实现中在训练前先用教师策略自行填充教师回放缓冲区详见原论文附录 A.5 Additional ablations for QDagger。教师规模原论文使用 DQN (Adam) 400M frames 的教师而 CleanRL 使用dqn_atari——即 DQN (Adam) 10M steps40M frames。Atari 预处理CleanRL 使用的是不使用 sticky action 的旧版 Atari 预处理而原论文使用 sticky action。JAX 版本qdagger_dqn_atari_jax_impalacnn.py功能特性使用 JAX、Flax 与 Optax 替代 PyTorchqdagger_dqn_atari_jax_impalacnn.py比 PyTorch 版本快约 25%–50%来自原文档的性能数据教师策略使用 Hugging Facecleanrl仓库中的dqn_atari_jax策略其余特性Atari 场景、Impala-CNN、预处理、观测/动作空间与 PyTorch 版本一致。安装与运行使用 uv 安装需同时安装 atari 与 jax 两个 extrauv pip install .[atari, 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 cleanrl/qdagger_dqn_atari_jax_impalacnn.py --env-id BreakoutNoFrameskip-v4 uv run python cleanrl/qdagger_dqn_atari_jax_impalacnn.py --env-id PongNoFrameskip-v4或使用 pippip install -r requirements/requirements-atari.txt pip install -r requirements/requirements-jax.txt pip install --upgrade jax[cuda]0.3.17 -f https://storage.googleapis.com/jax-releases/jax_cuda_releases.html python cleanrl/qdagger_dqn_atari_jax_impalacnn.py --env-id BreakoutNoFrameskip-v4 python cleanrl/qdagger_dqn_atari_jax_impalacnn.py --env-id PongNoFrameskip-v4注意JAX 不支持 Windows。官方安装文档建议使用 Windows Subsystem for Linux (WSL) 来安装 JAX即原文档中的 Windows 警告。与 PyTorch 版本的源码级差异网络定义改用 Flaxnn.linen训练状态使用扩展的TrainState在target_params中维护目标网络参数见 qdagger_dqn_atari_jax_impalacnn.py教师模型权重通过flax.serialization.from_bytes加载保存/评估用flax.serialization.to_bytes更新函数update用jax.jit编译蒸馏损失通过jax.vmap(kl_divergence_with_logits)对批量逐样本计算后取均值目标网络更新使用optax.incremental_update回放缓冲区放在 CPU 上devicecpu从缓冲区采样出的数据以.numpy()传入编译后的更新函数脚本开头设置了XLA_PYTHON_CLIENT_MEM_FRACTION0.7以控制 XLA 显存占用。其日志指标、实现细节与 PyTorch 版本一致可参考上文对应小节。实验与结果PyTorch 版本qdagger_dqn_atari_impalacnn.py以下为 10M steps40M frames下的平均回合回报环境qdagger_dqn_atari_impalacnn.py10M steps(40M frames)(Agarwal et al., 2022) 10M framesBreakoutNoFrameskip-v4295.55 ± 12.30275.15 ± 20.65PongNoFrameskip-v419.72 ± 0.20-BeamRiderNoFrameskip-v49284.99 ± 242.286514.25 ± 411.10学习曲线见 docs/rl-algorithms/qdagger/ 目录下的对应环境图片与dqn_atari的对比曲线见 docs/rl-algorithms/qdagger/compare.png。JAX 版本qdagger_dqn_atari_jax_impalacnn.py环境qdagger_dqn_atari_jax_impalacnn.py10M steps(40M frames)(Agarwal et al., 2022) 10M framesBreakoutNoFrameskip-v4335.08 ± 19.12275.15 ± 20.65PongNoFrameskip-v418.75 ± 0.19-BeamRiderNoFrameskip-v48024.75 ± 579.026514.25 ± 411.10JAX 版本的对比曲线见 docs/rl-algorithms/qdagger/jax/compare.png。从表格可见两个变体在 40M frames 预算下均达到或超过原始论文 10M frames 教师策略的成绩而 CleanRL 使用的教师规模40M frames远小于论文的 400M frames体现了复用先前计算在样本效率上的收益。需要注意的是以上对比的预算口径不同CleanRL 为 10M steps/40M frames论文为 10M frames不宜直接跨列比较绝对值。测试验证仓库中为两个变体均提供了端到端冒烟测试用极小的步数参数验证脚本可完整跑通tests/test_atari_gymnasium.pytest_qdagger_dqn_atari_impalacnn与test_qdagger_dqn_atari_impalacnn_eval后者额外加--save-model验证模型保存tests/test_atari_jax_gymnasium.pyJAX 版对应的两个测试。测试均以--learning-starts 10 --total-timesteps 16 --buffer-size 10 --batch-size 4 --teacher-steps 16 --offline-steps 16 --teacher-eval-episodes 1这类微型配置运行快速验证教师下载、数据收集、离线/在线训练与模型保存全链路。参考资源原始论文Agarwal, Rishabh, Max Schwarzer, Pablo Samuel Castro, Aaron Courville, and Marc G. Bellemare. Reincarnating Reinforcement Learning: Reusing Prior Computation to Accelerate Progress. arXiv, October 4, 2022arXiv:2206.01626官方参考实现google-research/reincarnating_rl原论文配套代码仅作参考背景本文档 docs/rl-algorithms/qdagger.md两个单文件实现cleanrl/qdagger_dqn_atari_impalacnn.py 与 cleanrl/qdagger_dqn_atari_jax_impalacnn.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),仅供参考