stable-baselines3 DQN 完全指南:从算法原理到 CartPole/Atari 实战配置
2026/9/15 19:22:45 网站建设 项目流程

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-baselines3(SB3)中 DQN 模块的算法思想、可用策略、全部构造参数、ε-greedy 探索与目标网络更新机制,并给出可直接运行的 CartPole 示例与 Atari 基准复现命令。读完本文,你将能独立完成 DQN 模型的训练、调参、保存加载与结果复现。

DQN 是什么:SB3 中基于 FQI 的深度 Q 学习实现

Deep Q Network(DQN)建立在 Fitted Q-Iteration(FQI)框架之上(相关论文可检索 arXiv:1312.5602 与 Nature 上的 Nature 14236 号文章),其核心思想是:用神经网络拟合动作价值函数 Q(s, a),并在迭代中稳定 Q 值估计。为了让神经网络训练真正收敛,原始 DQN 论文引入了三个关键技巧,这三个技巧也正是 SB3 中DQN类的实现基石:

  1. 经验回放(Replay Buffer):将与环境交互产生的转移(s, a, r, s')存入回放缓冲区,训练时随机采样小批量,打破样本间的时序相关性;
  2. 目标网络(Target Network):维护一份参数滞后更新的q_net_target,用它计算 TD 目标,避免"追逐移动靶"导致的训练发散;
  3. 梯度裁剪(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.py)
MultiInputPolicy字典(Dict)形式的多模态观测默认使用CombinedExtractor特征提取器(stable_baselines3/dqn/policies.py)

SB3 的 DQN 不支持循环策略(Recurrent policies),但支持多进程并行环境(Multi processing)。在 Gymnasium 空间支持方面,stable_baselines3/dqn/dqn.py 中supported_action_spaces=(spaces.Discrete,)明确限定了动作空间只能是 Discrete,而观测空间的兼容性如下:

空间类型动作(Action)观测(Observation)
Discrete✔️✔️
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_mode="human") # MlpPolicy 适用于向量观测;verbose=1 会在控制台输出设备与 wrapper 等信息 model = DQN("MlpPolicy", env, verbose=1) # total_timesteps 为总环境步数;log_interval=4 表示每 4 个 episode 打印一次日志 model.learn(total_timesteps=10000, log_interval=4) # 保存模型(会生成 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: # deterministic=True 表示不采样随机动作,而是直接取 argmax 的贪心动作 action, _states = model.predict(obs, deterministic=True) obs, reward, terminated, truncated, info = env.step(action) if terminated or truncated: obs, info = env.reset()

几点来自源码的补充说明:

  • model.predict(obs, deterministic=True)会走 stable_baselines3/dqn/dqn.py 中重写的predict:当deterministic=True时完全绕过 ε-greedy 探索,直接返回 Q 值最大的动作(QNetwork._predictq_values.argmax(dim=1),见 stable_baselines3/dqn/policies.py);
  • 官方提示:上述示例仅用于演示库的 API 用法,短时间训练出的智能体不一定能真正解决环境;若要获得良好性能,请使用 RL Zoo 中调优过的超参数(复现方法见下文"结果复现"一节)。

构造参数全解(默认值取自源码)

DQN的构造函数定义在 stable_baselines3/dqn/dqn.py,各参数及默认值、含义如下:

参数默认值说明
policy(必填)策略名或策略类,如"MlpPolicy""CnnPolicy""MultiInputPolicy"
env(必填)学习环境(Gym 注册名或环境实例;加载模型时可传None
learning_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回放缓冲区类(如HerReplayBuffer);None时自动选择
replay_buffer_kwargsNone传给回放缓冲区的额外关键字参数
optimize_memory_usageFalse启用省内存版回放缓冲区(约省一半内存,但实现更复杂)
n_steps1大于 1 时使用 n-step 回报(配合NStepReplayBuffer
target_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_archactivation_fnoptimizer_class等)
verbose0日志级别:0 无输出,1 输出设备/wrapper 等提示,2 输出调试信息
seedNone伪随机数种子
device"auto"运行设备(cpu/cuda/autoauto在可用时自动用 GPU)

learning_rate支持传入调度函数(Schedule),其入参是"剩余训练进度"(从 1 递减到 0),这在迁移学习或精细调参时非常有用。n_stepsgamma会在自动选择NStepReplayBuffer时被写入其构造参数(见 stable_baselines3/common/off_policy_algorithm.py)。

源码级原理:目标网络更新与 Polyak 更新

目标网络在策略构建时初始化:_build会创建两个独立的QNetworkq_netq_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_update(stable_baselines3/dqn/dqn.py):

self._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.py:target = (1 - tau) * target + tau * source,全部就地完成且处于no_grad上下文,不产生中间张量和计算图。当tau = 1.0(DQN 默认值)时等价于直接把在线网络参数硬拷贝到目标网络,这与原始 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()随机取动作,否则走策略贪心预测;注意deterministic=True时完全不探索,因此训练中应使用默认的deterministic=False,仅在评估/部署时置True

源码级原理:训练循环、Huber 损失与梯度裁剪

train()(stable_baselines3/dqn/dqn.py)是 DQN 的梯度更新核心,流程如下:

  1. 切换到训练模式,并按调度更新学习率;
  2. 从回放缓冲区采样一个小批量(replay_buffer.sample(batch_size));
  3. no_grad下用目标网络计算下一状态 Q 值并取max,构造 1 步 TD 目标:target_q = r + (1 - done) * gamma * max_a' Q_target(s', a')(n-step 时折扣为gamma**n_steps);
  4. th.gather取出在线网络对实际执行动作的 Q 值估计;
  5. 计算Huber 损失F.smooth_l1_loss,对离群点不敏感,比 MSE 更稳健);
  6. clip_grad_norm_max_grad_norm(默认 10)裁剪梯度范数后执行optimizer.step()

日志会记录train/n_updatestrain/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必须为 1,n-step 暂不支持字典观测);
  • n_steps > 1时 →NStepReplayBuffer(自动补充n_stepsgamma参数);
  • 否则 → 标准ReplayBuffer(实现见 stable_baselines3/common/buffers.py)。

ReplayBuffer采用环形存储(pos指针与full标志),采样时从[0, size)均匀随机抽取索引(stable_baselines3/common/buffers.py);buffer_size会按max(buffer_size // n_envs, 1)折算到单个环境。此外:

  • optimize_memory_usage=True会使用省内存变体,但再次learn时若缓冲区非空且reset_num_timesteps=True_setup_learn会截断最后一条轨迹并给出警告(stable_baselines3/common/off_policy_algorithm.py),此时应改用reset_num_timesteps=False
  • 若想结合 HER 做目标条件强化学习,可传入replay_buffer_class=HerReplayBuffer(此时构造时必须传入环境,见 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
  • 特征提取器默认值随策略而异:MlpPolicyFlattenExtractor(展平向量),CnnPolicyNatureCNN(Nature 论文风格卷积网络,此时net_arch自动置空),MultiInputPolicyCombinedExtractor
  • 图像输入默认normalize_images=True,即自动除以 255.0 归一化像素值。

这些默认值均可通过policy_kwargs覆盖,例如policy_kwargs=dict(net_arch=[128, 128], activation_fn=nn.Tanh, optimizer_class=th.optim.RMSprop)。测试用例 tests/test_run.py 展示了最小化配置:DQN("MlpPolicy", "CartPole-v1", policy_kwargs=dict(net_arch=[64, 64]), learning_starts=100, buffer_size=500, learning_rate=3e-4, verbose=1)即可完成一次可运行的训练。

模型保存与加载

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_env=True)。从源码可以推断两点约束:

  • train_freq"episode"单位要求env.num_envs == 1,多环境时只能用"step"(stable_baselines3/common/off_policy_algorithm.py);
  • n_envs > target_update_interval时,_setup_model会给出警告:目标网络将在每次env.step()后更新(对应n_envs步),相当于把更新间隔压到了下限 1(stable_baselines3/dqn/dqn.py),此时应适当调大target_update_interval

结果复现:Atari 基准与 RL Zoo 用法

官方文档在 docs/modules/dqn.md 中说明:DQN 在 Atari 游戏上的完整学习曲线可在对应 PR #110 中查看。要复现这些结果,需要借助 RL Zoo 仓库(SB3 官方超参数/基准仓库,注意该仓库为外部项目,clone 时使用其自身地址)。复现流程如下:

第 1 步:克隆 RL Zoo

git clone https://github.com/DLR-RM/rl-baselines3-zoo cd rl-baselines3-zoo/

第 2 步:运行基准训练(将$ENV_ID替换为环境 id,例如BreakoutNoFrameskip-v4):

python 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 DQN

RL Zoo 中为 DQN 调优过的超参数(含学习率、buffer_sizetarget_update_interval、探索退火参数等)是你在自己的环境上取得好成绩的最快起点。

常见问题与调参建议

结合上文源码分析,给出几条基于实现的实践建议:

  • 训练不收敛或发散:检查learning_starts是否太小(预热不足)、max_grad_norm是否过松、tau是否偏离 1.0 导致目标网络抖动;观察 TensorBoard 中的train/lossrollout/exploration_rate
  • 探索与利用失衡exploration_fraction决定 ε 退火快慢,稀疏奖励环境可适当增大初始探索时间;
  • 显存/内存不足:调小buffer_size或开启optimize_memory_usage=True(注意随之而来的轨迹截断警告);
  • 想要 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.md
  • DQN 相关测试(验证目标网络更新、探索、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),仅供参考

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

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

立即咨询