共计 4038 个字符,预计需要花费 11 分钟才能阅读完成。
背景与痛点:传统 Q 学习的困境
强化学习中的 Q 学习算法通过维护一个 Q 表格来存储状态 - 动作对应的价值,但在处理像 Atari 游戏这样的高维状态空间时(如 210×160 像素的屏幕画面),传统方法面临两个致命问题:

-
维度灾难 :对于连续或高维状态空间,Q 表格的存储需求呈指数级增长。以 Atari 的原始像素输入为例,可能的状态组合高达 256^(210×160×3) 种,根本无法存储。
-
特征提取困难:像素数据本身包含大量噪声和冗余信息,传统方法难以自动提取对决策有用的高层特征。
2013 年,DeepMind 团队首次提出用深度神经网络替代 Q 表格的想法,但直接结合面临两个主要挑战:
- 数据样本间强相关性导致训练不稳定
- 目标 Q 值随网络更新不断变化(类似于追逐移动的目标)
技术突破:两大核心机制
经验回放(Experience Replay)
传统强化学习按时间顺序处理样本,存在两个问题:
- 连续样本高度相关,导致梯度更新方向有偏
- 当前经验立即丢弃,数据利用率低
DQN 的解决方案是引入 经验回放缓冲区:
- 将智能体的经历(state, action, reward, next_state)存储在固定大小的循环缓冲区中
- 训练时随机采样一批历史经验用于更新网络
这样做的好处是:
- 打破样本相关性,使训练更稳定
- 可重复利用历史经验,提升数据效率
- 类似监督学习的 batch 训练方式
固定目标网络(Fixed Target Network)
Q 学习的目标值计算公式为:
target = r + γ * max(Q(s',a'))
如果使用同一个网络估算当前 Q 值和目标 Q 值,相当于在追逐一个不断变化的目标,容易导致训练震荡。
DQN 的解决方案是:
- 维护两个结构相同但参数独立的网络
- 主网络(online network):负责选择动作和实时更新
- 目标网络(target network):专门用于计算目标 Q 值
- 每隔 C 步将主网络参数复制到目标网络
这种延迟更新机制使目标值在一段时间内保持稳定,大幅提高了训练稳定性。
核心实现:PyTorch 代码详解
import torch
import torch.nn as nn
import torch.optim as optim
import random
from collections import deque
class DQN(nn.Module):
"""3 层卷积 + 2 层全连接的 Q 网络"""
def __init__(self, input_shape, n_actions):
super(DQN, self).__init__()
self.conv = nn.Sequential(nn.Conv2d(input_shape[0], 32, kernel_size=8, stride=4),
nn.ReLU(),
nn.Conv2d(32, 64, kernel_size=4, stride=2),
nn.ReLU(),
nn.Conv2d(64, 64, kernel_size=3, stride=1),
nn.ReLU())
conv_out_size = self._get_conv_out(input_shape)
self.fc = nn.Sequential(nn.Linear(conv_out_size, 512),
nn.ReLU(),
nn.Linear(512, n_actions)
)
def _get_conv_out(self, shape):
"""计算卷积层输出尺寸"""
o = self.conv(torch.zeros(1, *shape))
return int(torch.prod(torch.tensor(o.size())))
def forward(self, x):
conv_out = self.conv(x).view(x.size()[0], -1)
return self.fc(conv_out)
class DQNAgent:
def __init__(self, state_dim, action_dim, lr=1e-4, gamma=0.99):
self.gamma = gamma
self.replay_buffer = deque(maxlen=100000)
self.batch_size = 32
# 双网络结构
self.policy_net = DQN(state_dim, action_dim)
self.target_net = DQN(state_dim, action_dim)
self.target_net.load_state_dict(self.policy_net.state_dict())
self.target_net.eval() # 目标网络不计算梯度
self.optimizer = optim.Adam(self.policy_net.parameters(), lr=lr)
self.loss_fn = nn.MSELoss()
self.update_target_every = 1000 # 目标网络更新频率
def store_transition(self, state, action, reward, next_state, done):
"""存储经验到回放缓冲区"""
self.replay_buffer.append((state, action, reward, next_state, done))
def sample_batch(self):
"""随机采样一批经验"""
batch = random.sample(self.replay_buffer, self.batch_size)
states, actions, rewards, next_states, dones = zip(*batch)
return (torch.stack(states),
torch.tensor(actions),
torch.tensor(rewards),
torch.stack(next_states),
torch.tensor(dones, dtype=torch.float32)
)
def update_model(self, step_count):
"""更新主网络参数"""
if len(self.replay_buffer) < self.batch_size:
return
states, actions, rewards, next_states, dones = self.sample_batch()
# 计算当前 Q 值
current_q = self.policy_net(states).gather(1, actions.unsqueeze(1))
# 计算目标 Q 值(使用目标网络)with torch.no_grad():
next_q = self.target_net(next_states).max(1)[0]
target_q = rewards + (1 - dones) * self.gamma * next_q
# 计算损失并更新
loss = self.loss_fn(current_q.squeeze(), target_q)
self.optimizer.zero_grad()
loss.backward()
self.optimizer.step()
# 定期更新目标网络
if step_count % self.update_target_every == 0:
self.target_net.load_state_dict(self.policy_net.state_dict())
性能对比:Atari 游戏实验结果
在经典的 Atari 2600 游戏测试中,DQN 展现出显著优势:
| 游戏名称 | 人类平均分 | 传统 Q 学习 | DQN |
|---|---|---|---|
| Breakout | 31 | 无法运行 | 168 |
| Pong | -3 | -21 | 20 |
| Space Invaders | 1652 | 728 | 1976 |
关键发现:
- 在 23 款测试游戏中,DQN 在 12 款上超越人类专业玩家
- 相同训练步数下,DQN 得分普遍比传统方法高 5 -10 倍
- 学习到的特征具有可迁移性:在 Pong 上训练的网络可直接在类似游戏上表现良好
避坑指南:训练中的常见问题
问题 1:训练初期不收敛
现象:前几万步得分几乎无提升
原因:随机探索阶段,网络尚未学到有效策略
解决方案:
– 设置合理的 ε -greedy 衰减策略(如从 1.0 线性衰减到 0.1)
– 增加预热步数(warm-up steps),先纯探索积累经验
问题 2:训练后期震荡
现象:模型表现忽上忽下
原因:目标网络更新不及时导致过估计
解决方案:
– 降低目标网络更新频率(如从每 100 步改为每 1000 步)
– 尝试更软的更新方式:θ_target = τ*θ_policy + (1-τ)*θ_target(τ=0.01)
问题 3:梯度爆炸
现象:Loss 突然变为 NaN
解决方案:
– 添加梯度裁剪:torch.nn.utils.clip_grad_norm_(model.parameters(), 10)
– 减小学习率(推荐初始值 1e-4)
– 检查 reward 是否合理缩放(建议归一化到[-1,1])
进阶思考:DQN 的局限与改进
现有不足
- 过估计问题:max 操作会系统性高估 Q 值
- 效率问题:所有动作价值都需要重复计算
- 探索不足:ε-greedy 策略不够高效
改进方向
- Double DQN:
- 解耦动作选择和价值评估
- 使用主网络选择动作,目标网络评估价值
-
公式:
target = r + γ * Q_target(s', argmax(Q_policy(s'))) -
Dueling DQN:
- 网络结构分离状态价值和优势函数
- 公式:
Q(s,a) = V(s) + A(s,a) - mean(A(s,.)) -
能更高效学习哪些状态有价值,而与动作无关
-
Prioritized Experience Replay:
- 根据 TD 误差优先级采样
- 加速重要经验的学习
实践建议
推荐从简单环境开始实践:
- 先在 CartPole 等经典控制问题上验证代码正确性
- 使用 FrameStack 处理连续 4 帧画面作为状态输入
- 监控训练过程的关键指标:
- 平均回合奖励
- Q 值变化幅度
- 经验回放缓冲区中 TD 误差分布
完整实现可以参考 OpenAI Baselines 或 Stable-Baselines3 等开源库。期待看到你训练出的第一个游戏 AI!
