ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

stable-baselines3 DQN 完全指南:从算法原理到 CartPole/Atari 实战配置

stable-baselines3 DQN 完全指南:从算法原理到 CartPole/Atari 实战配置 stable-baselines3 DQN 完全指南从算法原理到 CartPole/Atari 实战配置【免费下载链接】stable-baselines3PyTorch version of Stable Baselines, reliable implementations of reinforcement learning algorithms.项目地址: https://gitcode.com/GitHub_Trending/st/stable-baselines3本指南以 docs/modules/dqn.md 为骨架结合 stable_baselines3/dqn/dqn.py 与 stable_baselines3/dqn/policies.py 的源码实现系统讲解 stable-baselines3SB3中 DQN 模块的算法思想、可用策略、全部构造参数、ε-greedy 探索与目标网络更新机制并给出可直接运行的 CartPole 示例与 Atari 基准复现命令。读完本文你将能独立完成 DQN 模型的训练、调参、保存加载与结果复现。DQN 是什么SB3 中基于 FQI 的深度 Q 学习实现Deep Q NetworkDQN建立在 Fitted Q-IterationFQI框架之上相关论文可检索 arXiv:1312.5602 与 Nature 上的 Nature 14236 号文章其核心思想是用神经网络拟合动作价值函数 Q(s, a)并在迭代中稳定 Q 值估计。为了让神经网络训练真正收敛原始 DQN 论文引入了三个关键技巧这三个技巧也正是 SB3 中DQN类的实现基石经验回放Replay Buffer将与环境交互产生的转移(s, a, r, s)存入回放缓冲区训练时随机采样小批量打破样本间的时序相关性目标网络Target Network维护一份参数滞后更新的q_net_target用它计算 TD 目标避免追逐移动靶导致的训练发散梯度裁剪Gradient Clipping通过max_grad_norm限制梯度的范数上限进一步提高稳定性。需要注意SB3 提供的只是 vanilla原始版Deep Q-Learning并不包含 Double-DQN、Dueling-DQN、Prioritized Experience Replay 等扩展这一点在 docs/modules/dqn.md 中有明确声明。如果你需要这些扩展可以借助社区维护的 sb3-contrib 等扩展仓库来实现。可用策略与空间支持矩阵DQN通过policy_aliases注册了三种开箱即用的策略见 stable_baselines3/dqn/dqn.py策略名适用观测类型底层实现MlpPolicy低维向量观测如 CartPole 的 4 维状态stable_baselines3/dqn/policies.py 中MlpPolicy DQNPolicy的别名CnnPolicy图像观测如 Atari 屏幕帧继承DQNPolicy默认使用NatureCNN特征提取器stable_baselines3/dqn/policies.pyMultiInputPolicy字典Dict形式的多模态观测默认使用CombinedExtractor特征提取器stable_baselines3/dqn/policies.pySB3 的 DQN 不支持循环策略Recurrent policies但支持多进程并行环境Multi processing。在 Gymnasium 空间支持方面stable_baselines3/dqn/dqn.py 中supported_action_spaces(spaces.Discrete,)明确限定了动作空间只能是 Discrete而观测空间的兼容性如下空间类型动作Action观测ObservationDiscrete✔️✔️Box❌✔️MultiDiscrete❌✔️MultiBinary❌✔️Dict❌✔️也就是说DQN 只能输出离散动作从有限动作集合中选一个但可以接收连续向量、图像乃至字典观测。若需处理连续动作空间请使用 SAC、TD3 或 PPO 等算法。快速上手CartPole-v1 训练示例官方文档给出的示例麻雀虽小五脏俱全覆盖了训练、保存、加载与推理全流程。这里对其逐行注释import gymnasium as gym from stable_baselines3 import DQN env gym.make(CartPole-v1, render_modehuman) # MlpPolicy 适用于向量观测verbose1 会在控制台输出设备与 wrapper 等信息 model DQN(MlpPolicy, env, verbose1) # total_timesteps 为总环境步数log_interval4 表示每 4 个 episode 打印一次日志 model.learn(total_timesteps10000, log_interval4) # 保存模型会生成 dqn_cartpole.zip 及附属数据文件 model.save(dqn_cartpole) # 删除变量以演示保存后加载 del model # remove to demonstrate saving and loading # 从文件加载模型也可传入新 env 继续训练 model DQN.load(dqn_cartpole) obs, info env.reset() while True: # deterministicTrue 表示不采样随机动作而是直接取 argmax 的贪心动作 action, _states model.predict(obs, deterministicTrue) obs, reward, terminated, truncated, info env.step(action) if terminated or truncated: obs, info env.reset()几点来自源码的补充说明model.predict(obs, deterministicTrue)会走 stable_baselines3/dqn/dqn.py 中重写的predict当deterministicTrue时完全绕过 ε-greedy 探索直接返回 Q 值最大的动作QNetwork._predict中q_values.argmax(dim1)见 stable_baselines3/dqn/policies.py官方提示上述示例仅用于演示库的 API 用法短时间训练出的智能体不一定能真正解决环境若要获得良好性能请使用 RL Zoo 中调优过的超参数复现方法见下文结果复现一节。构造参数全解默认值取自源码DQN的构造函数定义在 stable_baselines3/dqn/dqn.py各参数及默认值、含义如下参数默认值说明policy必填策略名或策略类如MlpPolicy、CnnPolicy、MultiInputPolicyenv必填学习环境Gym 注册名或环境实例加载模型时可传Nonelearning_rate1e-4优化器学习率可以是随训练进度递减的函数调度器注意此默认值来自 Stable Baselines 惯例而非 Nature 原文buffer_size1_000_000回放缓冲区容量可容纳的最大转移条数learning_starts100正式开始学习前先收集的随机探索步数预热期batch_size32每次梯度更新的小批量大小tau1.0目标网络软更新系数Polyak 更新1.0表示硬拷贝gamma0.99折扣因子train_freq4每多少步做一次训练支持(5, step)或(2, episode)元组写法gradient_steps1每次 rollout 后做多少步梯度更新-1表示与环境步数相同replay_buffer_classNone回放缓冲区类如HerReplayBufferNone时自动选择replay_buffer_kwargsNone传给回放缓冲区的额外关键字参数optimize_memory_usageFalse启用省内存版回放缓冲区约省一半内存但实现更复杂n_steps1大于 1 时使用 n-step 回报配合NStepReplayBuffertarget_update_interval10000每多少环境步更新一次目标网络exploration_fraction0.1整个训练过程中探索率ε线性退火所用的时间占比exploration_initial_eps1.0初始随机动作概率ε 初值exploration_final_eps0.05最终随机动作概率ε 终值max_grad_norm10梯度裁剪的范数上限stats_window_size100用于滚动日志统计平均回报、平均回合长度的 episode 窗口大小tensorboard_logNoneTensorBoard 日志目录None不记录policy_kwargsNone传给策略的额外参数如net_arch、activation_fn、optimizer_class等verbose0日志级别0 无输出1 输出设备/wrapper 等提示2 输出调试信息seedNone伪随机数种子deviceauto运行设备cpu/cuda/autoauto在可用时自动用 GPUlearning_rate支持传入调度函数Schedule其入参是剩余训练进度从 1 递减到 0这在迁移学习或精细调参时非常有用。n_steps与gamma会在自动选择NStepReplayBuffer时被写入其构造参数见 stable_baselines3/common/off_policy_algorithm.py。源码级原理目标网络更新与 Polyak 更新目标网络在策略构建时初始化_build会创建两个独立的QNetworkq_net与q_net_target并把在线网络权重拷贝到目标网络同时将目标网络置于评估模式set_training_mode(False)影响 BatchNorm/Dropout 行为见 stable_baselines3/dqn/policies.py。与 SAC/TD3 在train()内联更新目标网络不同DQN 的目标网络更新发生在环境交互阶段collect_rollouts()每步都会调用_on_step()其中按target_update_interval触发polyak_updatestable_baselines3/dqn/dqn.pyself._n_calls 1 # 多环境时按 n_envs 折算每 n_envs 步对应一次完整 env.step() if self._n_calls % max(self.target_update_interval // self.n_envs, 1) 0: polyak_update(self.q_net.parameters(), self.q_net_target.parameters(), self.tau) # 同步 BatchNorm 的 running 统计量对应 GH issue #996 polyak_update(self.batch_norm_stats, self.batch_norm_stats_target, 1.0)polyak_update的实现位于 stable_baselines3/common/utils.pytarget (1 - tau) * target tau * source全部就地完成且处于no_grad上下文不产生中间张量和计算图。当tau 1.0DQN 默认值时等价于直接把在线网络参数硬拷贝到目标网络这与原始 DQN 论文定期整体替换的设定一致若把tau调小如0.005则退化为每步缓慢追踪的软更新。源码级原理ε-greedy 探索与线性退火DQN 使用 ε-greedy 策略做探索以概率 ε 随机采样动作以概率 1-ε 取 Q 值最大的贪心动作。_setup_model中通过LinearSchedule构建探索率调度器stable_baselines3/dqn/dqn.py其数学定义见 stable_baselines3/common/utils.pyε 从exploration_initial_eps默认 1.0开始在前exploration_fraction默认 0.1即训练前 10% 的时间内线性退火到exploration_final_eps默认 0.05此后保持exploration_final_eps不变。每次环境步后_on_step()会依据_current_progress_remaining刷新self.exploration_rate并写入日志rollout/exploration_rate可通过 TensorBoard 观察退火曲线。实际采样逻辑在重写的predict中stable_baselines3/dqn/dqn.py非 deterministic 且np.random.rand() exploration_rate时从action_space.sample()随机取动作否则走策略贪心预测注意deterministicTrue时完全不探索因此训练中应使用默认的deterministicFalse仅在评估/部署时置True。源码级原理训练循环、Huber 损失与梯度裁剪train()stable_baselines3/dqn/dqn.py是 DQN 的梯度更新核心流程如下切换到训练模式并按调度更新学习率从回放缓冲区采样一个小批量replay_buffer.sample(batch_size)在no_grad下用目标网络计算下一状态 Q 值并取max构造 1 步 TD 目标target_q r (1 - done) * gamma * max_a Q_target(s, a)n-step 时折扣为gamma**n_steps用th.gather取出在线网络对实际执行动作的 Q 值估计计算Huber 损失F.smooth_l1_loss对离群点不敏感比 MSE 更稳健clip_grad_norm_按max_grad_norm默认 10裁剪梯度范数后执行optimizer.step()。日志会记录train/n_updates与train/loss。测试用例 tests/test_cnn.py 验证了梯度更新不改变目标网络、目标网络只被_on_step更新这一设计对 DQN 特设target_update_interval 1后手动调用_on_step()再执行train()断言目标网络参数发生变化而在线网络保持不变。经验回放机制与缓冲区自动选择DQN继承自 stable_baselines3/common/off_policy_algorithm.py 中的OffPolicyAlgorithm。构造模型时不传replay_buffer_class的话_setup_model会按以下规则自动选择stable_baselines3/common/off_policy_algorithm.py观测为Dict时 →DictReplayBuffer注意此时n_steps必须为 1n-step 暂不支持字典观测n_steps 1时 →NStepReplayBuffer自动补充n_steps与gamma参数否则 → 标准ReplayBuffer实现见 stable_baselines3/common/buffers.py。ReplayBuffer采用环形存储pos指针与full标志采样时从[0, size)均匀随机抽取索引stable_baselines3/common/buffers.pybuffer_size会按max(buffer_size // n_envs, 1)折算到单个环境。此外optimize_memory_usageTrue会使用省内存变体但再次learn时若缓冲区非空且reset_num_timestepsTrue_setup_learn会截断最后一条轨迹并给出警告stable_baselines3/common/off_policy_algorithm.py此时应改用reset_num_timestepsFalse若想结合 HER 做目标条件强化学习可传入replay_buffer_classHerReplayBuffer此时构造时必须传入环境见 stable_baselines3/common/off_policy_algorithm.py缓冲区可通过save_replay_buffer/load_replay_buffer单独保存与恢复stable_baselines3/common/off_policy_algorithm.py。learning_starts对应的预热期同样值得留意在达到该步数之前_sample_action一律随机采样动作stable_baselines3/common/off_policy_algorithm.py保证回放缓冲区内先积累足够多样化的转移避免冷启动训练崩溃。策略网络结构QNetwork 与特征提取器所有 DQN 策略都围绕QNetwork构建stable_baselines3/dqn/policies.py默认net_arch [64, 64]即观测经特征提取后接两个 64 维隐藏层输出维度为动作数action_space.n默认激活函数为nn.ReLU默认优化器为th.optim.Adam特征提取器默认值随策略而异MlpPolicy用FlattenExtractor展平向量CnnPolicy用NatureCNNNature 论文风格卷积网络此时net_arch自动置空MultiInputPolicy用CombinedExtractor图像输入默认normalize_imagesTrue即自动除以 255.0 归一化像素值。这些默认值均可通过policy_kwargs覆盖例如policy_kwargsdict(net_arch[128, 128], activation_fnnn.Tanh, optimizer_classth.optim.RMSprop)。测试用例 tests/test_run.py 展示了最小化配置DQN(MlpPolicy, CartPole-v1, policy_kwargsdict(net_arch[64, 64]), learning_starts100, buffer_size500, learning_rate3e-4, verbose1)即可完成一次可运行的训练。模型保存与加载model.save(dqn_cartpole)会序列化策略权重与优化器状态_get_torch_save_params返回[policy, policy.optimizer]而q_net/q_net_target两个别名会被排除_excluded_save_params见 stable_baselines3/dqn/dqn.py加载时由policy重建从而保证磁盘格式精简且版本兼容。DQN.load(dqn_cartpole)后可直接继续learn()或predict()若要换环境继续训练加载时传入env...即可。多环境训练注意事项DQN 支持多环境并行support_multi_envTrue。从源码可以推断两点约束train_freq的episode单位要求env.num_envs 1多环境时只能用stepstable_baselines3/common/off_policy_algorithm.py当n_envs target_update_interval时_setup_model会给出警告目标网络将在每次env.step()后更新对应n_envs步相当于把更新间隔压到了下限 1stable_baselines3/dqn/dqn.py此时应适当调大target_update_interval。结果复现Atari 基准与 RL Zoo 用法官方文档在 docs/modules/dqn.md 中说明DQN 在 Atari 游戏上的完整学习曲线可在对应 PR #110 中查看。要复现这些结果需要借助 RL Zoo 仓库SB3 官方超参数/基准仓库注意该仓库为外部项目clone 时使用其自身地址。复现流程如下第 1 步克隆 RL Zoogit clone https://github.com/DLR-RM/rl-baselines3-zoo cd rl-baselines3-zoo/第 2 步运行基准训练将$ENV_ID替换为环境 id例如BreakoutNoFrameskip-v4python train.py --algo dqn --env $ENV_ID --eval-episodes 10 --eval-freq 10000第 3 步绘制结果曲线python scripts/all_plots.py -a dqn -e Pong Breakout -f logs/ -o logs/dqn_results python scripts/plot_from_file.py -i logs/dqn_results.pkl -latex -l DQNRL Zoo 中为 DQN 调优过的超参数含学习率、buffer_size、target_update_interval、探索退火参数等是你在自己的环境上取得好成绩的最快起点。常见问题与调参建议结合上文源码分析给出几条基于实现的实践建议训练不收敛或发散检查learning_starts是否太小预热不足、max_grad_norm是否过松、tau是否偏离 1.0 导致目标网络抖动观察 TensorBoard 中的train/loss与rollout/exploration_rate探索与利用失衡exploration_fraction决定 ε 退火快慢稀疏奖励环境可适当增大初始探索时间显存/内存不足调小buffer_size或开启optimize_memory_usageTrue注意随之而来的轨迹截断警告想要 n-step 回报设置n_steps 1要求观测非 Dict只能处理离散动作若任务动作空间是Box连续空间请改用 SAC/TD3/PPO。进一步阅读仓库内资料算法总览与各算法对比docs/guide/algos.md自定义策略net_arch、特征提取器深入讲解docs/guide/custom_policy.md自定义环境与 Gymnasium 使用见 docs/guide/custom_env.md 与 docs/guide/vec_envs.md训练监控与 TensorBoard 用法docs/guide/tensorboard.mdDQN 相关测试验证目标网络更新、探索、CNN 支持tests/test_cnn.py、tests/test_run.py、tests/test_train_eval_mode.py回放缓冲区实现细节stable_baselines3/common/buffers.py离策略算法基类stable_baselines3/common/off_policy_algorithm.py【免费下载链接】stable-baselines3PyTorch version of Stable Baselines, reliable implementations of reinforcement learning algorithms.项目地址: https://gitcode.com/GitHub_Trending/st/stable-baselines3创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表