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类的实现基石:
- 经验回放(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.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._predict中q_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_rate | 1e-4 | 优化器学习率,可以是"随训练进度递减"的函数调度器;注意此默认值来自 Stable Baselines 惯例,而非 Nature 原文 |
buffer_size | 1_000_000 | 回放缓冲区容量(可容纳的最大转移条数) |
learning_starts | 100 | 正式开始学习前先收集的随机探索步数(预热期) |
batch_size | 32 | 每次梯度更新的小批量大小 |
tau | 1.0 | 目标网络软更新系数(Polyak 更新),1.0表示硬拷贝 |
gamma | 0.99 | 折扣因子 |
train_freq | 4 | 每多少步做一次训练,支持(5, "step")或(2, "episode")元组写法 |
gradient_steps | 1 | 每次 rollout 后做多少步梯度更新;-1表示与环境步数相同 |
replay_buffer_class | None | 回放缓冲区类(如HerReplayBuffer);None时自动选择 |
replay_buffer_kwargs | None | 传给回放缓冲区的额外关键字参数 |
optimize_memory_usage | False | 启用省内存版回放缓冲区(约省一半内存,但实现更复杂) |
n_steps | 1 | 大于 1 时使用 n-step 回报(配合NStepReplayBuffer) |
target_update_interval | 10000 | 每多少环境步更新一次目标网络 |
exploration_fraction | 0.1 | 整个训练过程中探索率(ε)线性退火所用的时间占比 |
exploration_initial_eps | 1.0 | 初始随机动作概率(ε 初值) |
exploration_final_eps | 0.05 | 最终随机动作概率(ε 终值) |
max_grad_norm | 10 | 梯度裁剪的范数上限 |
stats_window_size | 100 | 用于滚动日志统计(平均回报、平均回合长度)的 episode 窗口大小 |
tensorboard_log | None | TensorBoard 日志目录(None不记录) |
policy_kwargs | None | 传给策略的额外参数(如net_arch、activation_fn、optimizer_class等) |
verbose | 0 | 日志级别:0 无输出,1 输出设备/wrapper 等提示,2 输出调试信息 |
seed | None | 伪随机数种子 |
device | "auto" | 运行设备(cpu/cuda/auto,auto在可用时自动用 GPU) |
learning_rate支持传入调度函数(Schedule),其入参是"剩余训练进度"(从 1 递减到 0),这在迁移学习或精细调参时非常有用。n_steps与gamma会在自动选择NStepReplayBuffer时被写入其构造参数(见 stable_baselines3/common/off_policy_algorithm.py)。
源码级原理:目标网络更新与 Polyak 更新
目标网络在策略构建时初始化:_build会创建两个独立的QNetwork(q_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_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 的梯度更新核心,流程如下:
- 切换到训练模式,并按调度更新学习率;
- 从回放缓冲区采样一个小批量(
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必须为 1,n-step 暂不支持字典观测); n_steps > 1时 →NStepReplayBuffer(自动补充n_steps与gamma参数);- 否则 → 标准
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; - 特征提取器默认值随策略而异:
MlpPolicy用FlattenExtractor(展平向量),CnnPolicy用NatureCNN(Nature 论文风格卷积网络,此时net_arch自动置空),MultiInputPolicy用CombinedExtractor; - 图像输入默认
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 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_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),仅供参考