强化学习:DQN 算法原理与 PyTorch 实现(游戏 AI 场景)
强化学习: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^-$。这稳定了训练目标。
算法流程:
- 初始化主 Q 网络和目标 Q 网络(参数相同)。
- 初始化回放缓冲区。
- 对于每个 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$。
- 重复直到收敛。
在游戏 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):
- 修改状态处理:使用 CNN 替代全连接网络(例如
nn.Conv2d层处理图像)。 - 预处理图像:如灰度化、缩放和堆叠帧(例如 4 帧历史)。
- 调整超参数:增加缓冲区大小、使用更大的批次。
- 优化:添加帧跳过(frame skipping)减少计算量。
总结
DQN 算法通过结合深度学习和强化学习,在游戏 AI 中实现了从原始输入学习策略的能力。PyTorch 实现灵活且高效,但实际应用中需调参和优化(如使用 Double DQN 或 Prioritized Experience Replay)。在复杂游戏中,训练可能需要 GPU 加速。尝试运行代码后,您能观察到智能体在 CartPole 中的表现提升,验证了 DQN 的有效性。
更多推荐
所有评论(0)