深度解析2015 DeepMind DQN:从原理到实现深度强化学习的突破

1次阅读
没有评论

共计 4038 个字符,预计需要花费 11 分钟才能阅读完成。

image.webp

背景与痛点:传统 Q 学习的困境

强化学习中的 Q 学习算法通过维护一个 Q 表格来存储状态 - 动作对应的价值,但在处理像 Atari 游戏这样的高维状态空间时(如 210×160 像素的屏幕画面),传统方法面临两个致命问题:

深度解析 2015 DeepMind DQN:从原理到实现深度强化学习的突破

  1. 维度灾难 :对于连续或高维状态空间,Q 表格的存储需求呈指数级增长。以 Atari 的原始像素输入为例,可能的状态组合高达 256^(210×160×3) 种,根本无法存储。

  2. 特征提取困难:像素数据本身包含大量噪声和冗余信息,传统方法难以自动提取对决策有用的高层特征。

2013 年,DeepMind 团队首次提出用深度神经网络替代 Q 表格的想法,但直接结合面临两个主要挑战:

  • 数据样本间强相关性导致训练不稳定
  • 目标 Q 值随网络更新不断变化(类似于追逐移动的目标)

技术突破:两大核心机制

经验回放(Experience Replay)

传统强化学习按时间顺序处理样本,存在两个问题:

  1. 连续样本高度相关,导致梯度更新方向有偏
  2. 当前经验立即丢弃,数据利用率低

DQN 的解决方案是引入 经验回放缓冲区

  • 将智能体的经历(state, action, reward, next_state)存储在固定大小的循环缓冲区中
  • 训练时随机采样一批历史经验用于更新网络

这样做的好处是:

  1. 打破样本相关性,使训练更稳定
  2. 可重复利用历史经验,提升数据效率
  3. 类似监督学习的 batch 训练方式

固定目标网络(Fixed Target Network)

Q 学习的目标值计算公式为:

target = r + γ * max(Q(s',a'))

如果使用同一个网络估算当前 Q 值和目标 Q 值,相当于在追逐一个不断变化的目标,容易导致训练震荡。

DQN 的解决方案是:

  1. 维护两个结构相同但参数独立的网络
  2. 主网络(online network):负责选择动作和实时更新
  3. 目标网络(target network):专门用于计算目标 Q 值
  4. 每隔 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

关键发现:

  1. 在 23 款测试游戏中,DQN 在 12 款上超越人类专业玩家
  2. 相同训练步数下,DQN 得分普遍比传统方法高 5 -10 倍
  3. 学习到的特征具有可迁移性:在 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 的局限与改进

现有不足

  1. 过估计问题:max 操作会系统性高估 Q 值
  2. 效率问题:所有动作价值都需要重复计算
  3. 探索不足:ε-greedy 策略不够高效

改进方向

  1. Double DQN
  2. 解耦动作选择和价值评估
  3. 使用主网络选择动作,目标网络评估价值
  4. 公式:target = r + γ * Q_target(s', argmax(Q_policy(s')))

  5. Dueling DQN

  6. 网络结构分离状态价值和优势函数
  7. 公式:Q(s,a) = V(s) + A(s,a) - mean(A(s,.))
  8. 能更高效学习哪些状态有价值,而与动作无关

  9. Prioritized Experience Replay

  10. 根据 TD 误差优先级采样
  11. 加速重要经验的学习

实践建议

推荐从简单环境开始实践:

  1. 先在 CartPole 等经典控制问题上验证代码正确性
  2. 使用 FrameStack 处理连续 4 帧画面作为状态输入
  3. 监控训练过程的关键指标:
  4. 平均回合奖励
  5. Q 值变化幅度
  6. 经验回放缓冲区中 TD 误差分布

完整实现可以参考 OpenAI Baselines 或 Stable-Baselines3 等开源库。期待看到你训练出的第一个游戏 AI!

正文完
 0
评论(没有评论)