强化学习:DQN 算法原理与 PyTorch 实现(游戏 AI 场景)

Deep Q-Network(DQN)算法是强化学习中的一种先进方法,它将 Q-learning 与深度神经网络结合,特别适合处理高维状态空间(如游戏图像)。在游戏 AI 场景中(例如 Atari 游戏),DQN 能通过像素输入学习智能策略。下面我将逐步解释其原理,并提供 PyTorch 实现示例。实现基于简化环境(如 Gym 库的 CartPole),但原理可扩展到复杂游戏。

1. DQN 算法原理

DQN 的核心是学习一个 Q 函数,该函数估计在状态 $s$ 下执行动作 $a$ 的预期累积回报。Q-learning 的目标是最大化未来折扣回报: $$Q(s, a) = \mathbb{E}\left[ r + \gamma \max_{a'} Q(s', a') \mid s, a \right]$$ 其中 $r$ 是即时奖励,$\gamma$ 是折扣因子($0 < \gamma < 1$),$s'$ 是下一个状态。

DQN 的创新点在于使用深度神经网络(DNN)近似 Q 函数,并引入两个关键技术解决训练不稳定问题:

  • 经验回放(Experience Replay):存储经验元组 $(s, a, r, s', \text{done})$ 在回放缓冲区中,训练时随机采样一批数据。这打破了数据相关性,提高样本效率。
  • 目标网络(Target Network):使用一个独立的网络计算目标 Q 值,主网络参数 $\theta$ 定期同步到目标网络参数 $\theta^-$。这稳定了训练目标。

算法流程:

  1. 初始化主 Q 网络和目标 Q 网络(参数相同)。
  2. 初始化回放缓冲区。
  3. 对于每个 episode:
    • 初始化状态 $s$。
    • 对于每个步骤:
      • 使用 $\epsilon$-greedy 策略选择动作 $a$(以概率 $\epsilon$ 随机探索,否则选择 $\max_a Q(s, a)$)。
      • 执行 $a$,观察奖励 $r$ 和下一个状态 $s'$。
      • 存储 $(s, a, r, s', \text{done})$ 到缓冲区。
      • 从缓冲区采样一批经验。
      • 计算目标值 $y$: $$ y = \begin{cases} r & \text{if done} \ r + \gamma \max_{a'} Q_{\text{target}}(s', a'; \theta^-) & \text{otherwise} \end{cases} $$
      • 更新主网络:最小化均方误差损失 $L = \frac{1}{N} \sum (Q_{\text{main}}(s, a; \theta) - y)^2$。
      • 定期更新目标网络(例如每 C 步):$\theta^- \leftarrow \theta$。
  4. 重复直到收敛。

在游戏 AI 中,状态 $s$ 通常是图像帧,DQN 使用卷积神经网络(CNN)提取特征。优势包括处理原始像素的能力,但局限性是训练可能缓慢且对超参数敏感。

2. PyTorch 实现(基于 CartPole 环境)

这里使用 OpenAI Gym 的 CartPole-v1 环境作为简化游戏场景。状态是 4 维向量(非图像),但原理相同。代码包含完整训练循环。

前置依赖

确保安装 PyTorch 和 Gym:

pip install torch gym

完整代码
import gym
import random
import numpy as np
import torch
import torch.nn as nn
import torch.optim as optim
from collections import deque

# 定义 Q 网络(全连接网络)
class DQN(nn.Module):
    def __init__(self, state_dim, action_dim):
        super(DQN, self).__init__()
        self.fc1 = nn.Linear(state_dim, 128)
        self.fc2 = nn.Linear(128, 128)
        self.fc3 = nn.Linear(128, action_dim)
    
    def forward(self, x):
        x = torch.relu(self.fc1(x))
        x = torch.relu(self.fc2(x))
        return self.fc3(x)

# DQN 代理
class DQNAgent:
    def __init__(self, state_dim, action_dim):
        self.state_dim = state_dim
        self.action_dim = action_dim
        self.memory = deque(maxlen=10000)  # 回放缓冲区
        self.gamma = 0.99  # 折扣因子
        self.epsilon = 1.0  # 探索率
        self.epsilon_min = 0.01
        self.epsilon_decay = 0.995
        self.batch_size = 64
        self.update_target_every = 100  # 更新目标网络的步数
        
        self.main_network = DQN(state_dim, action_dim)
        self.target_network = DQN(state_dim, action_dim)
        self.target_network.load_state_dict(self.main_network.state_dict())  # 初始化相同
        self.optimizer = optim.Adam(self.main_network.parameters(), lr=0.001)
        self.loss_fn = nn.MSELoss()
        
        self.step_count = 0
    
    def remember(self, state, action, reward, next_state, done):
        self.memory.append((state, action, reward, next_state, done))
    
    def act(self, state):
        if np.random.rand() <= self.epsilon:
            return random.randrange(self.action_dim)  # 随机探索
        state = torch.FloatTensor(state).unsqueeze(0)
        with torch.no_grad():
            q_values = self.main_network(state)
        return torch.argmax(q_values).item()  # 选择最佳动作
    
    def replay(self):
        if len(self.memory) < self.batch_size:
            return
        
        batch = random.sample(self.memory, self.batch_size)
        states, actions, rewards, next_states, dones = zip(*batch)
        
        states = torch.FloatTensor(np.array(states))
        actions = torch.LongTensor(actions)
        rewards = torch.FloatTensor(rewards)
        next_states = torch.FloatTensor(np.array(next_states))
        dones = torch.BoolTensor(dones)
        
        # 计算当前 Q 值
        current_q = self.main_network(states).gather(1, actions.unsqueeze(1)).squeeze(1)
        
        # 计算目标 Q 值(使用目标网络)
        with torch.no_grad():
            next_q = self.target_network(next_states).max(1)[0]
        target_q = rewards + self.gamma * next_q * (~dones)  # 如果 done,则目标为 r
        
        # 更新主网络
        loss = self.loss_fn(current_q, target_q)
        self.optimizer.zero_grad()
        loss.backward()
        self.optimizer.step()
        
        # 衰减探索率
        if self.epsilon > self.epsilon_min:
            self.epsilon *= self.epsilon_decay
        
        # 定期更新目标网络
        self.step_count += 1
        if self.step_count % self.update_target_every == 0:
            self.target_network.load_state_dict(self.main_network.state_dict())

# 训练函数
def train_dqn(env, episodes=500):
    state_dim = env.observation_space.shape[0]
    action_dim = env.action_space.n
    agent = DQNAgent(state_dim, action_dim)
    
    for episode in range(episodes):
        state = env.reset()
        total_reward = 0
        done = False
        
        while not done:
            action = agent.act(state)
            next_state, reward, done, _ = env.step(action)
            agent.remember(state, action, reward, next_state, done)
            state = next_state
            total_reward += reward
            agent.replay()  # 训练网络
        
        print(f"Episode: {episode+1}, Total Reward: {total_reward}, Epsilon: {agent.epsilon:.3f}")
    
    return agent

# 主程序
if __name__ == "__main__":
    env = gym.make('CartPole-v1')
    agent = train_dqn(env)
    env.close()

代码说明
  • 环境设置:使用 CartPole-v1,状态是 4 维向量(位置、速度等),动作是 2 个(左/右)。奖励为 +1 每步,目标保持杆平衡。
  • Q 网络:全连接网络(3 层),输入状态维度,输出动作 Q 值。
  • 经验回放:缓冲区存储经验,采样时随机抽取批次。
  • 目标网络:每 100 步同步主网络参数。
  • 训练循环:$\epsilon$-greedy 策略逐渐减少探索($\epsilon$ 从 1.0 衰减到 0.01)。
  • 输出:每 episode 打印总奖励,训练后智能体能在 CartPole 中达到高分。
扩展到游戏 AI

对于图像输入的游戏(如 Atari):

  1. 修改状态处理:使用 CNN 替代全连接网络(例如 nn.Conv2d 层处理图像)。
  2. 预处理图像:如灰度化、缩放和堆叠帧(例如 4 帧历史)。
  3. 调整超参数:增加缓冲区大小、使用更大的批次。
  4. 优化:添加帧跳过(frame skipping)减少计算量。
总结

DQN 算法通过结合深度学习和强化学习,在游戏 AI 中实现了从原始输入学习策略的能力。PyTorch 实现灵活且高效,但实际应用中需调参和优化(如使用 Double DQN 或 Prioritized Experience Replay)。在复杂游戏中,训练可能需要 GPU 加速。尝试运行代码后,您能观察到智能体在 CartPole 中的表现提升,验证了 DQN 的有效性。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐