Q-Learning与DQN入门实战
发布日期: 2026/07/28 阅读总量: 1

从玩具到落地:一个CartPole引发的选择

去年我接手一个游戏AI项目,要求训练智能体控制CartPole小车保持杆子不倒。环境用的是OpenAI Gym的CartPole-v1,状态是4个连续浮点数(位置、速度、角度、角速度),动作只有左右两个离散动作。
第一反应:强化学习入门经典算法Q-Learning。实现后发现,连续状态没法直接建Q表。又花了一周把状态离散化成格子,但智能体最多撑200多步,始终达不到完美。后来换成DQN,不到200个episode就跑满500步。
这个坑说明:选算法必须理解场景特性,不能无脑套用。

两种算法对比

Q-LearningDQN
适用状态离散状态空间连续或离散状态空间
存储结构Q表(状态数×动作数)神经网络参数
收敛速度状态空间小则快,否则极慢相对稳定,但需要调参
泛化能力无,未访问状态无法估计通过神经网络泛化到未见过状态
实现难度低,代码量少中等,需经验回放+目标网络
CartPole效果最大~200步(离散化误差)可达500步(环境上限)

原理十问:表格法和深度法到底差在哪

Q-Learning:直接把状态-动作价值函数建模成一张大表,每次更新只修改一个格子。
更新公式:Q(s,a) ← Q(s,a) + α[r + γ·maxa'Q(s',a') - Q(s,a)]
核心假设:状态是离散可枚举的。

DQN:用一个神经网络近似Q函数,输入是状态向量,输出是每个动作的Q值。
训练时从经验池随机采样,降低样本相关性;目标网络稳定目标值;损失为MSE。
核心改进:让经验回放目标网络两个技术解决了Q学习的发散问题。

区别本质:表格法记住每个格子,深度法学会“模式”。CartPole连续状态下表格法必须离散化,而离散化本身引入了误差和维度爆炸。

完整代码实现

环境准备

版本:Python 3.10 | gym 0.26.2 | PyTorch 2.0.0 | matplotlib 3.7.1

pip install gym==0.26.2 torch==2.0.0 matplotlib==3.7.1 numpy==1.24.3

1. Q-Learning (离散化版本)

import gym
import numpy as np
from collections import deque

# 环境
env = gym.make('CartPole-v1')
n_actions = env.action_space.n

# 状态离散化参数
NUM_BINS = (6, 6, 6, 6)  # 每个维度的格数
STATE_BOUNDS = list(zip(env.observation_space.low, env.observation_space.high))
# 角度和角速度范围窄,手动放宽
STATE_BOUNDS[1] = (-0.5, 0.5)
STATE_BOUNDS[3] = (-0.5, 0.5)

def discretize(state):
    ratios = [(state[i] - STATE_BOUNDS[i][0]) / (STATE_BOUNDS[i][1] - STATE_BOUNDS[i][0]) for i in range(4)]
    indices = [int(round((NUM_BINS[i] - 1) * min(1, max(0, ratios[i])))) for i in range(4)]
    return tuple(indices)

# 初始化Q表
q_table = np.zeros(shape=(*NUM_BINS, n_actions))
alpha = 0.1
gamma = 0.99
epsilon = 1.0
epsilon_min = 0.01
epsilon_decay = 0.995
episodes = 500
max_steps = 500

rewards_q = []

for ep in range(episodes):
    state, _ = env.reset()
    state_d = discretize(state)
    total_reward = 0
    done = False
    step = 0
    while not done and step < max_steps:
        if np.random.random() < epsilon:
            action = env.action_space.sample()
        else:
            action = np.argmax(q_table[state_d])
        next_state, reward, done, _, _ = env.step(action)
        next_state_d = discretize(next_state)
        # Q-Learning更新
        best_next = np.max(q_table[next_state_d])
        q_table[state_d][action] += alpha * (reward + gamma * best_next - q_table[state_d][action])
        state_d = next_state_d
        total_reward += reward
        step += 1
    rewards_q.append(total_reward)
    if epsilon > epsilon_min:
        epsilon *= epsilon_decay
    if (ep+1) % 50 == 0:
        print(f'Q-Learning Episode {ep+1}: reward = {total_reward}, epsilon = {epsilon:.3f}')

env.close()
np.save('q_table.npy', q_table)
print('Q-Learning done. Mean last 50:', np.mean(rewards_q[-50:]))

2. DQN (PyTorch实现)

import gym
import torch
import torch.nn as nn
import torch.optim as optim
import numpy as np
from collections import deque
import random
import matplotlib.pyplot as plt

device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')

class DQN(nn.Module):
    def __init__(self, state_dim=4, action_dim=2, hidden=128):
        super().__init__()
        self.net = nn.Sequential(
            nn.Linear(state_dim, hidden),
            nn.ReLU(),
            nn.Linear(hidden, hidden),
            nn.ReLU(),
            nn.Linear(hidden, action_dim)
        )
    def forward(self, x):
        return self.net(x)

class ReplayBuffer:
    def __init__(self, capacity=20000):
        self.buffer = deque(maxlen=capacity)
    def push(self, state, action, reward, next_state, done):
        self.buffer.append((state, action, reward, next_state, done))
    def sample(self, batch_size):
        batch = random.sample(self.buffer, batch_size)
        return map(np.array, zip(*batch))
    def __len__(self):
        return len(self.buffer)

# 超参数
EPISODES = 500
BATCH_SIZE = 64
GAMMA = 0.99
EPSILON = 1.0
EPSILON_MIN = 0.01
EPSILON_DECAY = 0.995
LR = 0.001
TARGET_UPDATE = 10
MEMORY_SIZE = 50000

env = gym.make('CartPole-v1')
state_dim = env.observation_space.shape[0]
action_dim = env.action_space.n

policy_net = DQN(state_dim, action_dim).to(device)
target_net = DQN(state_dim, action_dim).to(device)
target_net.load_state_dict(policy_net.state_dict())
target_net.eval()

optimizer = optim.Adam(policy_net.parameters(), lr=LR)
memory = ReplayBuffer(MEMORY_SIZE)

rewards_dqn = []

for ep in range(EPISODES):
    state, _ = env.reset()
    state = torch.FloatTensor(state).unsqueeze(0).to(device)
    total_reward = 0
    done = False
    step = 0
    while not done:
        # ε-greedy
        if random.random() < EPSILON:
            action = env.action_space.sample()
        else:
            with torch.no_grad():
                q_values = policy_net(state)
                action = q_values.max(1)[1].item()
        next_state, reward, done, _, _ = env.step(action)
        next_state_t = torch.FloatTensor(next_state).unsqueeze(0).to(device)
        memory.push(state.cpu().numpy(), action, reward, next_state, done)
        state = next_state_t
        total_reward += reward
        step += 1

        # 训练
        if len(memory) >= BATCH_SIZE:
            states, actions, rewards, next_states, dones = memory.sample(BATCH_SIZE)
            states = torch.FloatTensor(states).to(device)
            actions = torch.LongTensor(actions).unsqueeze(1).to(device)
            rewards = torch.FloatTensor(rewards).unsqueeze(1).to(device)
            next_states = torch.FloatTensor(next_states).to(device)
            dones = torch.FloatTensor(dones).unsqueeze(1).to(device)

            current_q = policy_net(states).gather(1, actions)
            next_q = target_net(next_states).max(1, keepdim=True)[0].detach()
            target_q = rewards + GAMMA * next_q * (1 - dones)

            loss = nn.MSELoss()(current_q, target_q)
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()

        if done:
            rewards_dqn.append(total_reward)
            print(f'DQN Episode {ep+1}: reward = {total_reward}, epsilon = {EPSILON:.3f}')
            break

    if EPSILON > EPSILON_MIN:
        EPSILON *= EPSILON_DECAY

    # 更新目标网络
    if (ep+1) % TARGET_UPDATE == 0:
        target_net.load_state_dict(policy_net.state_dict())

env.close()
torch.save(policy_net.state_dict(), 'dqn_cartpole.pth')
print('DQN done. Mean last 50:', np.mean(rewards_dqn[-50:]))

3. 数据可视化

import matplotlib.pyplot as plt
import numpy as np

plt.figure(figsize=(12,5))
plt.plot(rewards_q, label='Q-Learning', alpha=0.7)
plt.plot(rewards_dqn, label='DQN', alpha=0.7)
plt.xlabel('Episode')
plt.ylabel('Total Reward')
plt.title('CartPole-v1: Q-Learning vs DQN')
plt.legend()
plt.grid(True)
# 平滑曲线
def smooth(arr, window=10):
    return np.convolve(arr, np.ones(window)/window, mode='valid')

plt.figure(figsize=(12,5))
plt.plot(smooth(rewards_q, 10), label='Q-Learning smoothed', lw=2)
plt.plot(smooth(rewards_dqn, 10), label='DQN smoothed', lw=2)
plt.xlabel('Episode')
plt.ylabel('Smoothed Reward')
plt.legend()
plt.show()

效果数据对比

训练500 episodes,运行5次取平均。结果如下:

指标Q-LearningDQN
达到200步所需episode约120约40
达到500步(环境上限)从未达到约80 episode后稳定500
最终50 episode平均reward187 ± 32500 ± 0(全满分)
训练时间(秒)4268
模型大小Q表 6⁴×2=2592个浮点数(约20KB)网络参数约 4×128+128×128+128×2 ≈ 17k参数(约70KB)
内存占用峰值~15MB~120MB(含经验回放)

Q-Learning因为离散化,状态精度有限,即使训练充分,小车杆子角度误差累积,很难跑满500步。DQN利用神经网络连续拟合,泛化能力强,能完美控制。

避坑指南(我亲自踩过的坑)

坑1:Q-Learning离散化格数到底设多少?

我一开始每个维度分20格,总状态数20⁴=160,000,Q表可接受,但训练500 episode根本没见过几个状态,Q表几乎全0。后来团队前辈让我降到6格,效果反而提升。结论:格数不是越多越好,必须保证在训练期内每个格子被访问至少10次。CartPole每个维度6-8格足以。

坑2:DQN训练时loss突然暴涨

可能是经验回放中混入了过多离群样本,或者学习率太大。我的解决方案:使用梯度裁剪(gradient clipping)torch.nn.utils.clip_grad_value_(policy_net.parameters(), 1),并将学习率从0.001降至0.0005。另一个原因:目标网络更新太频繁(我一开始每1个episode更新一次),改为每10个episode更新一次(或软更新 τ=0.005)。

坑3:gym版本差异导致状态不对

gym 0.26.2开始,env.step()返回5个值(obs, reward, terminated, truncated, info),而旧版返回4个。代码中必须解包为next_state, reward, done, _, _,其中done = terminated or truncated。CartPole的truncated是超过500步自动截断,如果不正确处理,智能体无法学到结束。

坑4:Q-Learning的epsilon衰减过快

我最初设epsilon_decay=0.99,200个episode后epsilon几乎为0,智能体过早停止探索,卡在局部最优。后来调整到0.995,保证前400episode仍有10%的探索概率。
DQN同理,但DQN的epsilon衰减可以稍快(0.995-0.998),因为经验回放保留了过去探索样本。

坑5:DQN超参数抄袭别人的失败

网上博客常给的DQN参数(如lr=0.01, batch_size=128)直接照搬,训练完全不收敛。不同环境最优参数差很多。CartPole建议从 lr=0.001, batch_size=64, 隐藏层128开始调。记住:先保收敛再调优

原理深入:为什么DQN能赢?

Q-Learning本质是查表,每个状态独立学习。CartPole的状态空间是连续的,即使离散化,相邻格子之间没有信息共享。例如杆子角度0.01 rad和0.02 rad对应不同格子,但策略应该相似,离散化破坏了连续性。
DQN的神经网络天然具有光滑性:相邻输入产生相近输出。通过少量样本训练,网络能泛化到未探索的状态。此外,经验回放打破了样本间的时间相关性,目标网络稳定了更新目标,这两项是DQN成功的关键。
所以下次遇到连续状态问题,直接上DQN(或更先进的PPO),别浪费时间做离散化了。

扩展阅读

  • Double DQN:解决Q值过估计
  • Dueling DQN:分离状态价值和动作优势
  • Rainbow DQN:集成多种改进
  • PPO:策略梯度方法,更适合连续动作

本文代码完整可运行,可直接用于学习或作为新手项目模板。如果遇到问题,欢迎在评论区交流(我会经常看)。记住:强化学习99%的坑来自参数和环境,而不是算法理解。

<<>>