用 PyTorch 在 CartPole 中训练 DQN

CartPole-v1 是一个离散动作控制任务:智能体在“向左移动小车”和“向右移动小车”之间选择,使车上的杆尽量保持直立。本例使用 Gymnasium 环境和 PyTorch,建立经验回放、Q 网络、目标网络以及完整训练循环。

每推进一个时间步,环境返回当前状态、奖励和结束信息。任务通常在杆偏离过大或小车距离中心超过 2.4 个单位时终止,每个有效时间步的奖励为 +1。因此,持续更长时间意味着累计更多奖励。Gymnasium CartPole 环境说明

当前示例直接使用环境返回的四个实数状态值,例如位置和速度,不进行缩放;把它们送入一个小型全连接网络,输出两个动作各自的预期回报。选择输出较大的动作即可形成贪心策略。

准备依赖和环境

需要已安装 PyTorch、Matplotlib,以及提供环境的 Gymnasium。Gymnasium 是原 Gym 项目的维护分支;当前代码使用 gymnasium as gym。

如果在 Colab 中,原文提供以下 Bash 单元。第一行是 notebook 单元魔法,不是普通终端命令:

%%bash
pip3 install gymnasium[classic_control]

普通终端中可以直接运行带引号的 pip3 install "gymnasium[classic_control]",前提是它对应你要运行代码的 Python 环境。本稿没有安装这些依赖或启动训练。

先导入模块、创建 CartPole 环境,并设置绘图和计算设备:

import gymnasium as gym
import math
import random
import matplotlib
import matplotlib.pyplot as plt
from collections import namedtuple, deque
from itertools import count

import torch
import torch.nn as nn
import torch.optim as optim
import torch.nn.functional as F

env = gym.make("CartPole-v1")
# set up matplotlib
is_ipython = 'inline' in matplotlib.get_backend()
if is_ipython:
    from IPython import display

plt.ion()

# if GPU is to be used
device = torch.device(
    "cuda" if torch.cuda.is_available() else
    "mps" if torch.backends.mps.is_available() else
    "cpu"
)
# To ensure reproducibility during training, you can fix the random seeds
# by uncommenting the lines below. This makes the results consistent across
# runs, which is helpful for debugging or comparing different approaches.
#
# That said, allowing randomness can be beneficial in practice, as it lets
# the model explore different training trajectories.
# seed = 42
# random.seed(seed)
# torch.manual_seed(seed)
# env.reset(seed=seed)
# env.action_space.seed(seed)
# env.observation_space.seed(seed)
# if torch.cuda.is_available():
#     torch.cuda.manual_seed(seed)

优先选择 CUDA,其次是可用的 MPS,否则使用 CPU。可选的随机种子代码原样保留为注释:调试和比较方法时可以固定种子,但单次固定种子的表现不能代表所有训练结果。

原文为网络运算、优化器和自动求导分别使用 torch.nn、torch.optim 和 PyTorch 的自动微分机制。代码不从画面读取状态,也不需要截取屏幕。

经验回放

经验回放保存智能体观察到的状态转移,让同一份交互数据能在稍后继续用于训练。随机抽样能够减少一个批次中连续转移的相关性;DQN 中,这有助于稳定训练。

需要两个结构:

  • Transition 保存 state、action、next_state 和 reward。
  • ReplayMemory 是有容量上限的循环缓冲区,保存近期转移,并能随机抽取一批样本。
Transition = namedtuple('Transition',
                        ('state', 'action', 'next_state', 'reward'))


class ReplayMemory(object):

    def __init__(self, capacity):
        self.memory = deque([], maxlen=capacity)

    def push(self, *args):
        """Save a transition"""
        self.memory.append(Transition(*args))

    def sample(self, batch_size):
        return random.sample(self.memory, batch_size)

    def __len__(self):
        return len(self.memory)

这里的状态是当前环境的四维向量。官方页面仍留有“屏幕差分图像”的旧文字;那与本页的当前代码不一致,不能用来解释这些张量。

DQN 的目标和更新公式

本例用确定性转移来简化公式。在随机环境中,相应表达还要对可能的状态转移取期望。

策略希望最大化折扣累计奖励:

R(t₀) = Σ[t=t₀…∞] γ^(t−t₀) · rₜ

对于有界奖励的无限时域,选择 0 ≤ γ < 1 使远期项衰减。较小的折扣更重视近期、较确定的奖励,也偏好更早获得同等大小的奖励。

如果已知最优动作价值函数 Q*(s,a),就可以选择:

π*(s) = argmaxₐ Q*(s,a)

实际不知道 Q*,于是用神经网络近似。某个策略的动作价值满足 Bellman 关系:

Qπ(s,a) = r + γ · Qπ(s′, π(s′))

DQN 的训练目标使用下一状态动作价值的最大值,时间差分误差写为:

δ = Q(s,a) − [r + γ · maxₐ′ Q(s′,a′)]

每次从回放缓冲区抽取一个批次,对批内误差的损失取平均。这里采用 Huber 损失:误差小时像平方误差,误差大时像绝对误差,可以减轻价值估计中离群值的影响。

误差范围 单样本损失
abs(δ) ≤ 1 δ² / 2
abs(δ) > 1 abs(δ) − 1/2

Q 网络

模型是一个前馈网络:输入维度来自环境状态长度,两个隐藏层各有 128 个单元,最后输出每个动作的价值。一次可以输入单个状态,也可以输入整个训练批次。

class DQN(nn.Module):

    def __init__(self, n_observations, n_actions):
        super(DQN, self).__init__()
        self.layer1 = nn.Linear(n_observations, 128)
        self.layer2 = nn.Linear(128, 128)
        self.layer3 = nn.Linear(128, n_actions)
    # Called with either one element to determine next action, or a batch
    # during optimization. Returns tensor([[left0exp,right0exp]...]).
    def forward(self, x):
        x = F.relu(self.layer1(x))
        x = F.relu(self.layer2(x))
        return self.layer3(x)

模型返回未经过 softmax 的价值估计。它预测的是采取动作后的预期回报,不是动作概率。输出列分别对应环境的两个动作。

超参数、探索策略和绘图

下面创建策略网络、目标网络与优化器,并定义两个辅助函数。

select_action 使用 ε-greedy:以一定概率均匀随机选动作,其他时候选当前网络认为价值最高的动作。随机探索概率从 EPS_START 指数衰减到 EPS_END;EPS_DECAY 越大,衰减越慢。

plot_durations 记录每回合持续的步数,并在有至少 100 个回合时绘制最近 100 回合的移动平均。Notebook 内会在每回合后更新图;普通脚本使用 Matplotlib 窗口。

# BATCH_SIZE is the number of transitions sampled from the replay buffer
# GAMMA is the discount factor as mentioned in the previous section
# EPS_START is the starting value of epsilon
# EPS_END is the final value of epsilon
# EPS_DECAY controls the rate of exponential decay of epsilon, higher means a slower decay
# TAU is the update rate of the target network
# LR is the learning rate of the ``AdamW`` optimizer
BATCH_SIZE = 128
GAMMA = 0.99
EPS_START = 0.9
EPS_END = 0.01
EPS_DECAY = 2500
TAU = 0.005
LR = 3e-4
# Get number of actions from gym action space
n_actions = env.action_space.n
# Get the number of state observations
state, info = env.reset()
n_observations = len(state)

policy_net = DQN(n_observations, n_actions).to(device)
target_net = DQN(n_observations, n_actions).to(device)
target_net.load_state_dict(policy_net.state_dict())

optimizer = optim.AdamW(policy_net.parameters(), lr=LR, amsgrad=True)
memory = ReplayMemory(10000)


steps_done = 0

def select_action(state):
    global steps_done
    sample = random.random()
    eps_threshold = EPS_END + (EPS_START - EPS_END) * \
        math.exp(-1. * steps_done / EPS_DECAY)
    steps_done += 1
    if sample > eps_threshold:
        with torch.no_grad():
            # t.max(1) will return the largest column value of each row.
            # second column on max result is index of where max element was
            # found, so we pick action with the larger expected reward.
            return policy_net(state).max(1).indices.view(1, 1)
    else:
        return torch.tensor([[env.action_space.sample()]], device=device, dtype=torch.long)

episode_durations = []

def plot_durations(show_result=False):
    plt.figure(1)
    durations_t = torch.tensor(episode_durations, dtype=torch.float)
    if show_result:
        plt.title('Result')
    else:
        plt.clf()
        plt.title('Training...')
    plt.xlabel('Episode')
    plt.ylabel('Duration')
    plt.plot(durations_t.numpy())
    # Take 100 episode averages and plot them too
    if len(durations_t) >= 100:
        means = durations_t.unfold(0, 100, 1).mean(1).view(-1)
        means = torch.cat((torch.zeros(99), means))
        plt.plot(means.numpy())
    plt.pause(0.001)  # pause a bit so that plots are updated
    if is_ipython:
        if not show_result:
            display.display(plt.gcf())
            display.clear_output(wait=True)
        else:
            display.display(plt.gcf())

本例参数含义:

参数 原例值 含义
BATCH_SIZE 128 每次从回放缓冲区抽样的转移数量
GAMMA 0.99 奖励折扣因子
EPS_START / EPS_END 0.9 / 0.01 随机探索的起始概率和渐近下限
EPS_DECAY 2500 探索概率指数衰减的尺度
TAU 0.005 目标网络的软更新速率
LR 0.0003 AdamW 学习率
回放容量 10000 最近转移的存储上限

目标网络一开始复制策略网络参数,之后逐步跟随。动作张量保持 torch.long,用于选择动作列;批状态为浮点张量。使用 torch.no_grad() 进行动作选择时,不为这一步建立反向传播图。

一步优化

optimize_model 先检查回放样本是否达到一个批次。抽样后,把“转移对象组成的列表”转成“各字段组成的批次”,再拼接状态、动作和奖励张量。

策略网络输出每个状态的两个价值,用 gather(1, action_batch) 取出实际采取的动作对应的价值。目标网络在无梯度模式下计算非终止下一状态的最大价值;真正终止的状态,下一状态价值设为 0。将它与当前奖励组成目标,再计算损失和梯度。

def optimize_model():
    if len(memory) < BATCH_SIZE:
        return
    transitions = memory.sample(BATCH_SIZE)
    # Transpose the batch (see https://stackoverflow.com/a/19343/3343043 for
    # detailed explanation). This converts batch-array of Transitions
    # to Transition of batch-arrays.
    batch = Transition(*zip(*transitions))
    # Compute a mask of non-final states and concatenate the batch elements
    # (a final state would've been the one after which simulation ended)
    non_final_mask = torch.tensor(tuple(map(lambda s: s is not None,
                                          batch.next_state)), device=device, dtype=torch.bool)
    non_final_next_states = torch.cat([s for s in batch.next_state
                                                if s is not None])
    state_batch = torch.cat(batch.state)
    action_batch = torch.cat(batch.action)
    reward_batch = torch.cat(batch.reward)
    # Compute Q(s_t, a) - the model computes Q(s_t), then we select the
    # columns of actions taken. These are the actions which would've been taken
    # for each batch state according to policy_net
    state_action_values = policy_net(state_batch).gather(1, action_batch)
    # Compute V(s_{t+1}) for all next states.
    # Expected values of actions for non_final_next_states are computed based
    # on the "older" target_net; selecting their best reward with max(1).values
    # This is merged based on the mask, such that we'll have either the expected
    # state value or 0 in case the state was final.
    next_state_values = torch.zeros(BATCH_SIZE, device=device)
    with torch.no_grad():
        next_state_values[non_final_mask] = target_net(non_final_next_states).max(1).values
    # Compute the expected Q values
    expected_state_action_values = (next_state_values * GAMMA) + reward_batch
    # Compute Huber loss
    criterion = nn.SmoothL1Loss()
    loss = criterion(state_action_values, expected_state_action_values.unsqueeze(1))

    # Optimize the model
    optimizer.zero_grad()
    loss.backward()
    # In-place gradient clipping
    torch.nn.utils.clip_grad_value_(policy_net.parameters(), 100)
    optimizer.step()

原例使用 SmoothL1Loss,与上面的阈值为 1 的 Huber 形式一致,并把梯度值截在给定范围后执行优化器更新。

这个原例有一个值得核对的边界:如果抽出的整个批次都是真正终止的转移,non_final_next_states 的列表可能为空,直接 torch.cat([]) 会失败。上面的原代码原样保留;迁移到其他环境时,可为这个分支加保护。例如,用下面的两段替换原函数中相应的构造和赋值,其他语句保持原顺序:

next_states = [s for s in batch.next_state if s is not None]
non_final_next_states = torch.cat(next_states) if next_states else None

# 保留原函数中 state_batch、action_batch、reward_batch、
# state_action_values 的计算,再执行下面的部分。
next_state_values = torch.zeros(BATCH_SIZE, device=device)
with torch.no_grad():
    if non_final_next_states is not None:
        next_state_values[non_final_mask] = target_net(
            non_final_next_states
        ).max(1).values

这是额外的编辑建议,不是官方原例的运行结果。所有批次均终止时,下一状态价值保持为零。代码资产同时保存原版和仅增加这一保护的版本,修改文件带有明确说明。

完整训练循环

每个回合先重置环境,把初始四维状态变成带批维度的张量。每一步选择动作、与环境交互、记录奖励、存入回放,再执行一次优化和目标网络软更新。

软更新为:

θ′ ← τ · θ + (1−τ) · θ′

if torch.cuda.is_available() or torch.backends.mps.is_available():
    num_episodes = 600
else:
    num_episodes = 50
for i_episode in range(num_episodes):
    # Initialize the environment and get its state
    state, info = env.reset()
    state = torch.tensor(state, dtype=torch.float32, device=device).unsqueeze(0)
    for t in count():
        action = select_action(state)
        observation, reward, terminated, truncated, _ = env.step(action.item())
        reward = torch.tensor([reward], device=device)
        done = terminated or truncated
        if terminated:
            next_state = None
        else:
            next_state = torch.tensor(observation, dtype=torch.float32, device=device).unsqueeze(0)

        # Store the transition in memory
        memory.push(state, action, next_state, reward)

        # Move to the next state
        state = next_state

        # Perform one step of the optimization (on the policy network)
        optimize_model()
        # Soft update of the target network's weights
        # θ′ ← τ θ + (1 −τ )θ′
        target_net_state_dict = target_net.state_dict()
        policy_net_state_dict = policy_net.state_dict()
        for key in policy_net_state_dict:
            target_net_state_dict[key] = policy_net_state_dict[key]*TAU + target_net_state_dict[key]*(1-TAU)
        target_net.load_state_dict(target_net_state_dict)
        if done:
            episode_durations.append(t + 1)
            plot_durations()
            break

print('Complete')
plot_durations(show_result=True)
plt.ioff()
plt.show()

terminated 与 truncated 需要分开理解:

  • 两者任何一个为真,都结束当前回合。
  • 只有真正的终止 terminated 才把 next_state 设为 None,使该转移不继续估计后续价值。
  • 仅因时间限制而截断时,代码保留下一状态,继续在训练目标中使用它的价值估计。

有 CUDA 或 MPS 时,原例训练 600 回合;CPU 路径训练 50 回合,主要为了让教程运行时间较短。原文提醒,50 回合不足以观察到好的 CartPole 表现,并给出较长训练可达到 500 步的预期。这是教程中的经验说明,不能当作每次运行都会达到的保证。本稿没有执行训练、没有生成时长曲线,也没有测量收敛结果。

查看自己的训练曲线时,还应区分单回合的噪声和移动平均;如果要比较策略,应使用多种随机种子和独立评估回合,而不是仅依赖一条训练轨迹。

把数据流串起来

官方最后的数据流图表达了下面的闭环:

  1. 策略网络或随机探索选择动作。
  2. Gymnasium 环境执行动作,产生下一状态、奖励和结束标志。
  3. 把转移保存到经验回放,并在每次交互后进行优化。
  4. 从回放随机抽取批次,用策略网络计算当前动作价值,用较旧的目标网络计算训练目标。
  5. 更新策略网络,并在每一步软更新目标网络。

原图的教学关系已在这里用文字完整保留。图片方案会绘制概念示意图,不生成虚假的训练曲线或调试器截图。

来源与许可

作者为 Adam Paszke、Mark Towers 及 PyTorch 教程贡献者。原文为 Reinforcement Learning (DQN) Tutorial,完整源文件位于官方教程仓库。原页面标注更新日期为 2025 年 6 月 16 日、上次核验为 2024 年 11 月 5 日;本稿核对日期为 2026 年 10 月 3 日。

本文按 BSD 3-Clause 再利用,完整保留原例结构与代码,正文汉化并纠正旧屏幕输入措辞;另加全终止批次保护建议与未实测说明。代码资产包含原版、修改版及许可证。本稿未进行环境交互、GPU/CPU 训练或性能测试。

BSD 3-Clause 许可全文
BSD 3-Clause License

Copyright (c) 2017-2022, Pytorch contributors
All rights reserved.

Redistribution and use in source and binary forms, with or without
modification, are permitted provided that the following conditions are met:

* Redistributions of source code must retain the above copyright notice, this
list of conditions and the following disclaimer.
* Redistributions in binary form must reproduce the above copyright notice,
this list of conditions and the following disclaimer in the documentation
and/or other materials provided with the distribution.

* Neither the name of the copyright holder nor the names of its
contributors may be used to endorse or promote products derived from
this software without specific prior written permission.
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS “AS IS”
AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.

© 版权声明
THE END
喜欢就支持一下吧
点赞0 分享
评论 抢沙发

请登录后发表评论

    暂无评论内容