评价此页

循环 DQN:训练循环策略#

创建日期:2023年11月08日 | 最后更新:2025年01月27日 | 最后验证:未验证

作者Vincent Moens

您将学到什么
  • 如何在 TorchRL 的 actor 中加入 RNN

  • 如何将基于记忆的策略与经验回放池(replay buffer)和损失模块结合使用

先决条件
  • PyTorch v2.0.0

  • gym[mujoco]

  • tqdm

概述#

基于记忆的策略不仅在观测值部分可观测时至关重要,而且在需要考虑时间维度以做出明智决策时也同样关键。

循环神经网络长期以来一直是基于记忆的策略的常用工具。其核心思想是在两个连续步骤之间保持循环状态(recurrent state)的记忆,并将其与当前观测值一起作为策略的输入。

本教程展示了如何使用 TorchRL 将 RNN 合并到策略中。

主要学习内容

  • 在 TorchRL 的 actor 中加入 RNN;

  • 将该基于记忆的策略与经验回放池和损失模块配合使用。

在 TorchRL 中使用 RNN 的核心思路是利用 TensorDict 作为数据载体,在步骤间传递隐藏状态。我们将构建一个策略,从当前的 TensorDict 中读取先前的循环状态,并将当前的循环状态写入下一个状态的 TensorDict 中。

Data collection with a recurrent policy

如图所示,环境使用零初始化的循环状态填充 TensorDict,策略将其与观测值一起读取以产生动作,并输出将用于下一步的循环状态。当调用 step_mdp() 函数时,来自下一个状态的循环状态会被带入当前的 TensorDict。让我们看看这是如何在实践中实现的。

如果您在 Google Colab 中运行此代码,请确保安装以下依赖项

!pip3 install torchrl
!pip3 install gym[mujoco]
!pip3 install tqdm

设置#

import torch
import tqdm
from tensordict.nn import TensorDictModule as Mod, TensorDictSequential as Seq
from torch import nn
from torchrl.collectors import SyncDataCollector
from torchrl.data import LazyMemmapStorage, TensorDictReplayBuffer
from torchrl.envs import (
    Compose,
    ExplorationType,
    GrayScale,
    InitTracker,
    ObservationNorm,
    Resize,
    RewardScaling,
    set_exploration_type,
    StepCounter,
    ToTensorImage,
    TransformedEnv,
)
from torchrl.envs.libs.gym import GymEnv
from torchrl.modules import ConvNet, EGreedyModule, LSTMModule, MLP, QValueModule
from torchrl.objectives import DQNLoss, SoftUpdate

is_fork = multiprocessing.get_start_method() == "fork"
device = (
    torch.device(0)
    if torch.cuda.is_available() and not is_fork
    else torch.device("cpu")
)

环境#

照例,第一步是构建我们的环境:它有助于我们定义问题并据此构建策略网络。在本教程中,我们将运行一个基于像素的单实例 CartPole gym 环境,并添加一些自定义转换:转为灰度图、调整大小为 84x84、缩小奖励值以及标准化观测值。

注意

StepCounter 转换是辅助性的。由于 CartPole 任务的目标是尽可能延长轨迹,计数步骤可以帮助我们跟踪策略的性能。

对于本教程,有两个转换非常重要:

  • InitTracker 将通过在 TensorDict 中添加一个 "is_init" 布尔掩码来标记对 reset() 的调用,该掩码将跟踪哪些步骤需要重置 RNN 隐藏状态。

  • TensorDictPrimer 转换技术性更强。使用 RNN 策略并不强制要求它。然而,它告知环境(以及随后的收集器)需要预留一些额外的键。一旦添加,对 env.reset() 的调用将使用零张量填充 primer 中指定的条目。因为知道策略需要这些张量,收集器会在收集期间传递它们。最终,我们将把隐藏状态存储在经验回放池中,这将帮助我们在损失模块中引导 RNN 计算(否则这些计算将以 0 初始化)。总之:不包含此转换不会严重影响策略的训练,但会导致循环键在收集的数据和经验回放池中消失,这反过来会导致训练效果略逊一筹。幸运的是,我们提供的 LSTMModule 配备了一个辅助方法来帮我们构建该转换,所以我们可以等构建它时再处理!

env = TransformedEnv(
    GymEnv("CartPole-v1", from_pixels=True, device=device),
    Compose(
        ToTensorImage(),
        GrayScale(),
        Resize(84, 84),
        StepCounter(),
        InitTracker(),
        RewardScaling(loc=0.0, scale=0.1),
        ObservationNorm(standard_normal=True, in_keys=["pixels"]),
    ),
)

一如既往,我们需要手动初始化我们的标准化常量。

env.transform[-1].init_stats(1000, reduce_dim=[0, 1, 2], cat_dim=0, keep_dims=[0])
td = env.reset()

策略 (Policy)#

我们的策略将包含 3 个组件:一个 ConvNet 主干、一个 LSTMModule 记忆层,以及一个将 LSTM 输出映射到动作值的浅层 MLP 块。

卷积网络#

我们构建了一个带有 torch.nn.AdaptiveAvgPool2d 的卷积网络,它将输出压缩为一个大小为 64 的向量。ConvNet 可以辅助我们完成此操作。

feature = Mod(
    ConvNet(
        num_cells=[32, 32, 64],
        squeeze_output=True,
        aggregator_class=nn.AdaptiveAvgPool2d,
        aggregator_kwargs={"output_size": (1, 1)},
        device=device,
    ),
    in_keys=["pixels"],
    out_keys=["embed"],
)

我们在数据批次上执行第一个模块,以获取输出向量的大小。

n_cells = feature(env.reset())["embed"].shape[-1]

LSTM 模块#

TorchRL 提供了一个专门的 LSTMModule 类,用于在代码库中合并 LSTM。它是 TensorDictModuleBase 的子类:因此,它拥有一组 in_keysout_keys,指明了在模块执行期间应该读取和写入/更新哪些值。该类带有这些属性的可自定义预定义值,以方便构建。

注意

使用限制:该类支持几乎所有 LSTM 功能,例如 dropout 或多层 LSTM。然而,为了遵循 TorchRL 的约定,此 LSTM 必须将 batch_first 属性设置为 True,这在 PyTorch 中并非默认值。不过,我们的 LSTMModule 改变了此默认行为,因此我们可以直接进行原生调用。

此外,LSTM 不能将 bidirectional 属性设置为 True,因为这在在线设置中无法使用。在这种情况下,默认值是正确的。

lstm = LSTMModule(
    input_size=n_cells,
    hidden_size=128,
    device=device,
    in_key="embed",
    out_key="embed",
)

让我们看看 LSTM Module 类,特别是它的 in_keys 和 out_keys。

print("in_keys", lstm.in_keys)
print("out_keys", lstm.out_keys)

我们可以看到,这些值包含我们指定为 in_key(和 out_key)的键,以及循环键名称。out_keys 前面带有“next”前缀,表示它们需要写入“下一个” TensorDict 中。我们使用此约定(可以通过传递 in_keys/out_keys 参数覆盖)来确保对 step_mdp() 的调用将循环状态移动到根 TensorDict 中,使其在下一次调用时可供 RNN 使用(参见引言中的图)。

如前所述,我们还可以向环境中添加一个可选转换,以确保循环状态传递到缓冲区。make_tensordict_primer() 方法正是执行此操作。

env.append_transform(lstm.make_tensordict_primer())

就是这样!在添加了 primer 之后,我们可以打印环境以检查一切是否正常。

print(env)

MLP#

我们使用单层 MLP 来表示我们将用于策略的动作值。

mlp = MLP(
    out_features=2,
    num_cells=[
        64,
    ],
    device=device,
)

并将偏置填充为零。

mlp[-1].bias.data.fill_(0.0)
mlp = Mod(mlp, in_keys=["embed"], out_keys=["action_value"])

使用 Q 值选择动作#

策略的最后一部分是 Q 值模块。Q 值模块 QValueModule 将读取由我们的 MLP 产生的 "action_values" 键,并从中获取最大值的动作。我们唯一需要做的是指定动作空间,这可以通过传递字符串或动作规格(action-spec)来完成。这允许我们使用 Categorical(有时称为“稀疏”)编码或其独热(one-hot)版本。

qval = QValueModule(spec=env.action_spec)

注意

TorchRL 还提供了一个包装类 torchrl.modules.QValueActor,它将一个模块与 QValueModule 一起包装在一个 Sequential 中,正如我们在这里明确做的那样。这样做几乎没有什么优势,且过程透明度较低,但最终结果将与我们这里所做的相似。

现在,我们可以将所有内容放在一个 TensorDictSequential 中。

stoch_policy = Seq(feature, lstm, mlp, qval)

由于 DQN 是一种确定性算法,探索是其中的关键部分。我们将使用 epsilon 为 0.2 并逐渐衰减至 0 的 \(\epsilon\)-greedy 策略。此衰减通过调用 step() 实现(见下方的训练循环)。

exploration_module = EGreedyModule(
    annealing_num_steps=1_000_000, spec=env.action_spec, eps_init=0.2
)
stoch_policy = Seq(
    stoch_policy,
    exploration_module,
)

将模型用于损失计算#

我们构建的模型非常适合在序列设置中使用。然而,类 torch.nn.LSTM 可以使用 cuDNN 优化的后端来更快地在 GPU 设备上运行 RNN 序列。我们不想错过这样一个加快训练循环的机会!要使用它,我们只需告诉 LSTM 模块在被损失函数使用时以“循环模式”(recurrent-mode)运行。由于我们通常需要 LSTM 模块的两个副本,我们通过调用 set_recurrent_mode() 方法来实现,该方法将返回一个新的 LSTM 实例(具有共享权重),该实例将假设输入数据本质上是序列性的。

policy = Seq(feature, lstm.set_recurrent_mode(True), mlp, qval)

因为我们还有一些未初始化的参数,我们在创建优化器等之前应该初始化它们。

policy(env.reset())

DQN 损失#

我们的 DQN 损失需要传递策略,并再次传递动作空间。虽然这看起来多余,但很重要,因为我们希望确保 DQNLossQValueModule 类是兼容的,但又不是强依赖于彼此。

为了使用 Double-DQN,我们要求使用 delay_value 参数,这将创建一个网络参数的不可微副本,用作目标网络。

loss_fn = DQNLoss(policy, action_space=env.action_spec, delay_value=True)

由于我们使用的是双 DQN,我们需要更新目标参数。我们将使用一个 SoftUpdate 实例来完成这项工作。

updater = SoftUpdate(loss_fn, eps=0.95)

optim = torch.optim.Adam(policy.parameters(), lr=3e-4)

收集器和经验回放池#

我们构建最简单的数据收集器。我们将尝试用一百万帧来训练我们的算法,每次扩展 50 帧的缓冲区。该缓冲区被设计为存储 2 万条轨迹,每条轨迹 50 步。在每个优化步骤(每次数据收集 16 次)中,我们将从缓冲区中收集 4 个项目,总共 200 个转换。我们将使用 LazyMemmapStorage 存储将数据保留在磁盘上。

注意

为了效率起见,我们在这里只运行了几千次迭代。在实际设置中,总帧数应设置为 100 万。

collector = SyncDataCollector(env, stoch_policy, frames_per_batch=50, total_frames=200, device=device)
rb = TensorDictReplayBuffer(
    storage=LazyMemmapStorage(20_000), batch_size=4, prefetch=10
)

训练循环#

为了跟踪进度,我们将每 50 次数据收集在环境中运行一次策略,并在训练后绘制结果。

utd = 16
pbar = tqdm.tqdm(total=1_000_000)
longest = 0

traj_lens = []
for i, data in enumerate(collector):
    if i == 0:
        print(
            "Let us print the first batch of data.\nPay attention to the key names "
            "which will reflect what can be found in this data structure, in particular: "
            "the output of the QValueModule (action_values, action and chosen_action_value),"
            "the 'is_init' key that will tell us if a step is initial or not, and the "
            "recurrent_state keys.\n",
            data,
        )
    pbar.update(data.numel())
    # it is important to pass data that is not flattened
    rb.extend(data.unsqueeze(0).to_tensordict().cpu())
    for _ in range(utd):
        s = rb.sample().to(device, non_blocking=True)
        loss_vals = loss_fn(s)
        loss_vals["loss"].backward()
        optim.step()
        optim.zero_grad()
    longest = max(longest, data["step_count"].max().item())
    pbar.set_description(
        f"steps: {longest}, loss_val: {loss_vals['loss'].item(): 4.4f}, action_spread: {data['action'].sum(0)}"
    )
    exploration_module.step(data.numel())
    updater.step()

    with set_exploration_type(ExplorationType.DETERMINISTIC), torch.no_grad():
        rollout = env.rollout(10000, stoch_policy)
        traj_lens.append(rollout.get(("next", "step_count")).max().item())

让我们绘制结果。

if traj_lens:
    from matplotlib import pyplot as plt

    plt.plot(traj_lens)
    plt.xlabel("Test collection")
    plt.title("Test trajectory lengths")

结论#

我们已经了解了如何将 RNN 合并到 TorchRL 的策略中。你现在应该能够:

  • 创建一个充当 TensorDictModule 的 LSTM 模块。

  • 通过 InitTracker 转换向 LSTM 模块指示何时需要重置。

  • 将此模块合并到策略和损失模块中。

  • 确保收集器意识到循环状态条目,以便它们能够与其余数据一起存储在经验回放池中。

进一步阅读#

  • TorchRL 文档可以在这里找到。