Deep Q-Network(DQN)
用神经网络近似高维状态的 Q 值,并以 Replay 与 Target Network 稳定训练。
学习目标
- 定义Deep Q-Network(DQN)并复述输入输出。
- 从公式计算一个最小数值例子。
- 说明它与相邻算法的区别、失败模式和适用场景。
为什么重要
Deep Q-Network(DQN)位于“状态—行动—反馈—更新”学习链中的关键位置。
掌握它能帮助学习者判断算法使用的数据、策略归属和稳定性边界。
学习前需要掌握
背景与问题
强化学习面对序贯决策:动作会改变之后能看到的状态和奖励,因此样本通常并非独立同分布。
同一算法的效果取决于环境、探索策略、函数近似、随机种子和评测协议,单次曲线不足以下结论。
概念定义
DQN 输入状态,输出每个离散动作的 Q 值。
本页使用“问题定义—数学目标—更新过程—代码—失败诊断”的顺序组织,避免只记算法缩写。
直观理解
不再维护巨大 Q 表,而用网络从状态特征预测所有动作分数。
把价值看作“未来累计收益的估计”,把策略看作“在状态下如何选动作的规则”;算法差异主要在估计谁、使用谁生成的数据以及如何更新。
核心原理
Replay 打乱相关样本。
Target Network 延缓目标漂移。
Online Network 选择/拟合当前 Q。
DQN loss 常用 MSE 或 Huber。
数学表达
DQN loss
Target 参数 θ⁻ 在一段时间内固定。
变量含义
G_t从时刻 t 开始的折扣累计回报。R_{t+1}执行动作后收到的下一步奖励。γ折扣因子,通常位于 [0,1)。
计算步骤
- 1计算 bootstrap 目标:1+0.9×2=2.8。
- 2目标与当前估计差为 2.8−0.5=2.3。
- 3若学习率 α=0.1,新估计为 0.5+0.1×2.3=0.73。
代码对应位置
示例 1PyTorch DQN:Online、Target 与 Replay 的最小训练步:完整展示批次张量、DQN target、Huber loss、反向传播和 Target Network 同步。
变量解释
| 变量 | 含义 |
|---|---|
G_t | 从时刻 t 开始的折扣累计回报。 |
R_{t+1} | 执行动作后收到的下一步奖励。 |
γ | 折扣因子,通常位于 [0,1)。 |
完整数值示例
Deep Q-Network(DQN)的最小计算
已知条件
- 即时奖励为 1
- 下一状态估计为 2
- 折扣因子 γ=0.9
- 当前估计为 0.5
- 1
计算 bootstrap 目标:1+0.9×2=2.8。
- 2
目标与当前估计差为 2.8−0.5=2.3。
- 3
若学习率 α=0.1,新估计为 0.5+0.1×2.3=0.73。
一次更新后估计从 0.5 变为 0.73;是否收敛需要持续采样与满足相应条件。
处理前后对比
- 较大 α 更新快但噪声和震荡更强。
- 较大 γ 更重视远期奖励,但误差传播范围更长。
- 使用真实回报与 bootstrap 目标会带来不同偏差—方差权衡。
计算与实现步骤
- 1
与环境交互并写 Replay
- 2
采样 mini-batch
- 3
Online 计算 Q(s,a)
- 4
Target 计算 y
- 5
优化 TD loss
- 6
周期同步 Target
代码实现
示例 1
PyTorch DQN:Online、Target 与 Replay 的最小训练步
example_01.py用途:完整展示批次张量、DQN target、Huber loss、反向传播和 Target Network 同步。
import random
from collections import deque
import torch
from torch import nn
torch.manual_seed(7)
online = nn.Sequential(nn.Linear(4, 32), nn.ReLU(), nn.Linear(32, 2))
target = nn.Sequential(nn.Linear(4, 32), nn.ReLU(), nn.Linear(32, 2))
target.load_state_dict(online.state_dict())
target.eval()
optimizer = torch.optim.Adam(online.parameters(), lr=1e-3)
replay = deque(maxlen=1000)
for i in range(64):
state = torch.tensor([i % 4 == j for j in range(4)], dtype=torch.float32)
next_state = torch.roll(state, 1)
replay.append((state, i % 2, 1.0 if i % 7 == 0 else 0.0, next_state, False))
batch = random.Random(7).sample(list(replay), 32)
states = torch.stack([x[0] for x in batch])
actions = torch.tensor([x[1] for x in batch]).unsqueeze(1)
rewards = torch.tensor([x[2] for x in batch])
next_states = torch.stack([x[3] for x in batch])
dones = torch.tensor([x[4] for x in batch], dtype=torch.float32)
q_sa = online(states).gather(1, actions).squeeze(1)
with torch.no_grad():
next_q = target(next_states).max(1).values
td_target = rewards + 0.99 * (1 - dones) * next_q
loss = nn.SmoothL1Loss()(q_sa, td_target)
optimizer.zero_grad()
loss.backward()
nn.utils.clip_grad_norm_(online.parameters(), 10.0)
optimizer.step()
target.load_state_dict(online.state_dict())
print('loss:', round(loss.item(), 4))代码解析
解析始终位于完整代码下方,并按实际代码段逐项对应。
输入数据与任务
完整展示批次张量、DQN target、Huber loss、反向传播和 Target Network 同步。
loss- 当前预测与目标之间的损失值。
Step 1 · 1–4 行
导入当前步骤需要的数值计算、预处理、模型或评价工具。依赖集中写在代码开头,便于复现。
import random
from collections import deque
import torch
from torch import nnStep 2 · 6–12 行
Online Network 产生 Q(s,a),Target Network 产生相对稳定的 TD 目标。
torch.manual_seed(7)
online = nn.Sequential(nn.Linear(4, 32), nn.ReLU(), nn.Linear(32, 2))
target = nn.Sequential(nn.Linear(4, 32), nn.ReLU(), nn.Linear(32, 2))
target.load_state_dict(online.state_dict())
target.eval()
optimizer = torch.optim.Adam(online.parameters(), lr=1e-3)
replay = deque(maxlen=1000)Step 3 · 14–17 行
SmoothL1Loss、梯度裁剪和周期性 Target 同步共同改善稳定性。
for i in range(64):
state = torch.tensor([i % 4 == j for j in range(4)], dtype=torch.float32)
next_state = torch.roll(state, 1)
replay.append((state, i % 2, 1.0 if i % 7 == 0 else 0.0, next_state, False))Step 4 · 19–24 行
执行当前代码段,并把得到的状态传给下一步。
batch = random.Random(7).sample(list(replay), 32)
states = torch.stack([x[0] for x in batch])
actions = torch.tensor([x[1] for x in batch]).unsqueeze(1)
rewards = torch.tensor([x[2] for x in batch])
next_states = torch.stack([x[3] for x in batch])
dones = torch.tensor([x[4] for x in batch], dtype=torch.float32)Step 5 · 26–36 行
计算损失对参数的梯度,对应数学表达中的 ∇J;梯度只给出局部变化方向。
q_sa = online(states).gather(1, actions).squeeze(1)
with torch.no_grad():
next_q = target(next_states).max(1).values
td_target = rewards + 0.99 * (1 - dones) * next_q
loss = nn.SmoothL1Loss()(q_sa, td_target)
optimizer.zero_grad()
loss.backward()
nn.utils.clip_grad_norm_(online.parameters(), 10.0)
optimizer.step()
target.load_state_dict(online.state_dict())
print('loss:', round(loss.item(), 4))代码与数学原理
参数更新由当前梯度与学习率共同决定。
平方误差代码对应均方损失。
预期输出或运行结果
打印一个有限的正损失值;具体数值由固定初始化与批次决定。
常见错误 · 3 条
- 只报告最好的一次随机种子
- 训练回报与评测回报混用
- 终止状态仍错误 bootstrap
实际应用
- 序贯决策
- 控制与资源分配
常见错误
文档来源
强化学习可靠资料扩展资料
外部原始教材或论文- 相关定义、公式与算法章节
扩展内容说明
未找到可直接映射的本地强化学习文档;中文直觉、数值例子、代码和交互演示属于扩展解释,算法定义与公式以所列教材或原论文为依据。
算法属性与数据边界
失败模式、风险与性能
失败模式
- 探索不足
- 目标漂移
- 函数近似不稳定
目标与安全风险
- 奖励函数与真实目标不一致会诱发奖励黑客。
- 部署策略的行动权限必须由环境和应用层约束。
性能与复现
- 样本效率、墙钟时间和显存占用需要分别报告。
- 应使用多个随机种子、置信区间和固定评测策略。
常见问题
Deep Q-Network(DQN)是 on-policy 还是 off-policy?
本页算法/方法按 off-policy 组织。
网页是否会训练模型?
不会。所有图表使用固定种子或解析公式在浏览器本地计算,不执行页面中的示例代码。
资料来源与核对日期
核对日期:2026-08-30。算法定义与公式依据以下外部教材或原论文;中文直觉、教学代码、对照与部署建议属于本站扩展解释。
推荐学习资料
Hugging Face Deep Reinforcement Learning Course
从 Q-Learning、DQN、Policy Gradient 和 Actor-Critic 逐步进入 PPO 与多智能体。
Hugging Face · Hugging Face
资料笔记
Playing Atari with Deep Reinforcement Learning
提出使用深度网络、经验回放和目标网络从像素输入学习 Atari 控制策略。
Mnih et al. · arXiv
资料笔记
vwxyzjn/cleanrl
以单文件实现和实验报告呈现 DQN、PPO、SAC 等深度强化学习算法。
CleanRL Contributors · GitHub
仓库信息
阅读可复现的单文件 RL 算法实现
主要语言:Python包含数据集或实验需要额外环境