PyTorch强化学习实战——基于图像-文本融合的自动化网页导航

PyTorch强化学习实战------基于图像-文本融合的自动化网页导航

    • [0. 前言](#0. 前言)
    • [1. 在网页导航智能体中引入文本数据](#1. 在网页导航智能体中引入文本数据)
    • [2. 代码实现](#2. 代码实现)
    • [3. 运行结果](#3. 运行结果)
    • 相关链接

0. 前言

我们已经学习了强化学习在网页导航与浏览器自动化中的应用,介绍了 MiniWoB++ 基准测试,该环境提供像素观测、文本描述和 DOM 元素等多模态输入,动作空间涵盖鼠标键盘操作。并且实现了基于异步优势演员-评论家 (Asynchronous Advantage Actor-Critic, A3C) 算法的按钮点击智能体,将动作空间简化为网格化点击,通过卷积网络处理图像观测。实验结果表明该智能体能解决简单任务(如点击对话框),但在处理多步骤、依赖文本描述或违反马尔可夫性质的任务时效果不佳。

在本节中,我们将改进点击智能体,将把问题文本描述纳入模型。我们知道,某些任务包含关键信息(如需要点击的标签页索引或待勾选条目列表),这些信息通过文本描述提供。虽然相同信息会显示在图像观测顶部,但像素并非简单文本的最佳表征形式。

1. 在网页导航智能体中引入文本数据

为处理文本数据,需将模型输入从单一图像扩展为图像与文本的组合。借鉴 textworld 一节中的文本处理经验,循环神经网络 (Recurrent Neural Network, RNN)是自然选择(对此类小问题或许非最优解,但具备灵活性与扩展性)。

2. 代码实现

本节将重点介绍实现的关键要点,完整代码参见 wob_click_mm_train.py 模块。相比点击模型,文本扩展的改动并不大。

首先需要让 MiniWoBClickWrapper 保留来自观测的文本。为了保留文本,我们需要在包装器构造函数中传递 keep_text=True,该类将返回包含 NumPy 数组和文本字符串的元组,而非单独的图像数组组。随后需调整模型以处理此类元组而非批量数组------这涉及两个环节:智能体(使用模型选择动作时)和训练代码。

为以模型友好方式适配观测,可使用 preprocessor 函数:其核心思想是通过可调用函数将观测列表转换为可直接传入模型的格式。默认 preprocessorNumPy 数组列表转为 PyTorch 张量(可选复制到 GPU 内存),但本节需要更复杂的转换------将图像打包为张量的同时特殊处理文本字符串。此时可重写默认预处理器并传入 Agent 类。

(1) 理论上得益于 PyTorch 的灵活性,预处理器功能可移至模型内部,但默认预处理器在观测为纯 NumPy 数组时能大幅简化流程。以下是 model.py 模块中的 preprocessor 类源代码:

python 复制代码
MM_EMBEDDINGS_DIM = 50
MM_HIDDEN_SIZE = 128
MM_MAX_DICT_SIZE = 100

TOKEN_UNK = "#unk"

class MultimodalPreprocessor:
    log = logging.getLogger("MulitmodalPreprocessor")

    def __init__(self, max_dict_size: int = MM_MAX_DICT_SIZE,
                 device: torch.device = torch.device('cpu')):
        self.max_dict_size = max_dict_size
        self.token_to_id = {TOKEN_UNK: 0}
        self.next_id = 1
        self.tokenizer = TweetTokenizer(preserve_case=True)
        self.device = device

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

在上述代码的构造函数中,我们创建了从词元到标识符的映射(该映射将动态扩展),并使用 nltk 包创建了分词器。

(2) 接下来是 __call__() 方法,该方法负责转换批次数据:

python 复制代码
    def __call__(self, batch: tt.Tuple[tt.Any, ...] | tt.List[tt.Tuple[tt.Any, ...]]):
        tokens_batch = []

        if isinstance(batch, tuple):
            batch_iter = zip(*batch)
        else:
            batch_iter = batch
        for img_obs, txt_obs in batch_iter:
            tokens = self.tokenizer.tokenize(txt_obs)
            idx_obs = self.tokens_to_idx(tokens)
            tokens_batch.append((img_obs, idx_obs))
        # sort batch decreasing to seq len
        tokens_batch.sort(key=lambda p: len(p[1]), reverse=True)
        img_batch, seq_batch = zip(*tokens_batch)
        lens = list(map(len, seq_batch))

preprocessor 的目标是将(图像,文本)元组的批次转换为两个对象:第一个必须是形状为 (batch_size, 3, 210, 160) 的图像数据张量,第二个必须是以打包序列形式存储的文本描述词元批次。打包序列是 PyTorch 的一种数据结构,适用于RNN 高效处理

实际上,批次可能呈现两种形式:可能是包含图像批次和文本批次的元组,也可能是由单个(图像,文本词元)样本组成的元组列表。这种差异源于 VectorEnvgym.Tuple 观测空间的不同处理方式。不过这些细节在此并不重要,我们只需通过检查批次变量类型并执行必要处理来应对差异。

转换的第一步是对文本字符串进行分词,并将每个词元转换为整数 ID 列表。随后按词元长度递减排序批次------这是底层 cuDNN 库实现高效 RNN 处理的要求。

然后,我们将图像转换为张量,将序列转换为填充序列(即批次大小×最长序列长度的矩阵):

python 复制代码
        # convert data into the target form
        # images
        img_v = torch.FloatTensor(np.asarray(img_batch)).to(self.device)
        # sequences
        seq_arr = np.zeros(
            shape=(len(seq_batch), max(len(seq_batch[0]), 1)), dtype=np.int64)
        for idx, seq in enumerate(seq_batch):
            seq_arr[idx, :len(seq)] = seq
            # Map empty sequences into single #UNK token
            if len(seq) == 0:
                lens[idx] = 1
        seq_v = torch.LongTensor(seq_arr).to(self.device)
        seq_p = rnn_utils.pack_padded_sequence(seq_v, lens, batch_first=True)
        return img_v, seq_p

(3) tokens_to_idx() 函数将词元列表转换为 ID 列表:

python 复制代码
    def tokens_to_idx(self, tokens):
        res = []
        for token in tokens:
            idx = self.token_to_id.get(token)
            if idx is None:
                if self.next_id == self.max_dict_size:
                    self.log.warning("Maximum size of dict reached, token "
                                     "'%s' converted to #UNK token", token)
                    idx = 0
                else:
                    idx = self.next_id
                    self.next_id += 1
                    self.token_to_id[token] = idx
            res.append(idx)
        return res

问题在于,我们无法预先知道文本描述中词典的大小。一种方案是在字符级别处理,将单个字符输入 RNN,但这会导致序列过长难以处理。另一种方案是硬编码合理的词典大小(例如 100 个词符),并为未见过的新词符动态分配 ID。本实现采用后者,但这种方法可能不适用于文本描述包含随机生成字符串的 MiniWoB 问题。潜在解决方案包括使用字符级分词或预定义词典。

(4) 接下来,实现模型类:

python 复制代码
class ModelMultimodal(nn.Module):
    def __init__(self, input_shape: tt.Tuple[int, ...], n_actions: int,
                 max_dict_size: int = MM_MAX_DICT_SIZE):
        super(ModelMultimodal, self).__init__()

        self.conv = nn.Sequential(
            nn.Conv2d(input_shape[0], 64, 5, stride=5),
            nn.ReLU(),
            nn.Conv2d(64, 64, 3, stride=2),
            nn.ReLU(),
            nn.Flatten(),
        )
        size = self.conv(torch.zeros(1, *input_shape)).size()[-1]

        self.emb = nn.Embedding(max_dict_size, MM_EMBEDDINGS_DIM)
        self.rnn = nn.LSTM(MM_EMBEDDINGS_DIM, MM_HIDDEN_SIZE, batch_first=True)
        self.policy = nn.Linear(size + MM_HIDDEN_SIZE*2, n_actions)
        self.value = nn.Linear(size + MM_HIDDEN_SIZE*2, 1)

主要差异在于新增的嵌入层(将整型词元 ID 转换为稠密词符向量)和长短期记忆 (Long Short-Term Memory, LSTM) RNN。卷积层与 RNN 层的输出被拼接后馈入策略头和价值头,因此它们的输入维度是图像与文本特征的组合。

(5) _concat_features() 函数将图像特征和 RNN 特征拼接为单个张量:

python 复制代码
    def _concat_features(self, img_out: torch.Tensor,
                         rnn_hidden: torch.Tensor | tt.Tuple[torch.Tensor, ...]):
        batch_size = img_out.size()[0]
        if isinstance(rnn_hidden, tuple):
            flat_h = list(map(lambda t: t.view(batch_size, -1), rnn_hidden))
            rnn_h = torch.cat(flat_h, dim=1)
        else:
            rnn_h = rnn_hidden.view(batch_size, -1)
        return torch.cat((img_out, rnn_h), dim=1)

(6) 最后,在 forward() 函数中,我们期望 preprocessor 准备好的两个对象:一个包含输入图像的张量和一个批次的打包序列:

python 复制代码
    def forward(self, x: tt.Tuple[torch.Tensor, rnn_utils.PackedSequence]):
        x_img, x_text = x

        # deal with text data
        emb_out = self.emb(x_text.data)
        emb_out_seq = rnn_utils.PackedSequence(emb_out, x_text.batch_sizes)
        rnn_out, rnn_h = self.rnn(emb_out_seq)

        # extract image features
        xx = x_img / 255.0
        conv_out = self.conv(xx)
        feats = self._concat_features(conv_out, rnn_h)
        return self.policy(feats), self.value(feats)

图像数据通过卷积层处理,文本数据经 RNN 处理;随后将结果拼接并计算策略和价值输出。

以上就是新增代码的主要部分。训练脚本 wob_click_mm_train.py 基本上是 wob_click_train.py 的副本,仅对包装器创建、模型和预处理器进行了小幅修改。

3. 运行结果

在点击按钮环境中进行了多次实验,该环境目标是在多个随机按钮中进行选择。下图展示了该环境的几种情况:

如下图所示,过 3 小时训练后,模型已学会点击操作(每回合平均步数降至 5-7 步)并获得 0.2 的平均奖励。但后续训练未见明显改善,这可能表明需要调整超参数,或源于环境本身的模糊性。该环境有时会显示多个相同标题的按钮,但仅有一个能提供正奖励。上图第三部分的示例中就存在两个完全相同的"Submit"按钮。。

另一个文本描述至关重要的环境是点击标签页 (click-tab),该环境要求智能体点击随机选定的特定标签页。具体界面如下图所示:

在这个环境中,训练并未成功,这略显反常,因为任务看起来比点击按钮任务更简单(点击位置是固定的),我们可以通过调整超参数来尝试解决此环境。

相关链接

PyTorch强化学习实战(1)------强化学习(Reinforcement Learning,RL)详解

PyTorch强化学习实战(2)------强化学习环境库Gymnasium

PyTorch强化学习实战(3)------Gymnasium API扩展功能

PyTorch强化学习实战(4)------PyTorch基础

PyTorch强化学习实战(5)------PyTorch Ignite 事件驱动机制与实践

PyTorch强化学习实战(6)------交叉熵方法详解与实现

PyTorch强化学习实战(7)------表格学习与贝尔曼方程

PyTorch强化学习实战(8)------Q学习详解与实现

PyTorch强化学习实战(9)------深度Q学习

PyTorch强化学习实战(10)------强化学习高级组件

PyTorch强化学习实战(11)------N步DQN(N-step DQN)

PyTorch强化学习实战(12)------Double DQN(DDQN)

PyTorch强化学习实战(13)------噪声网络(NoisyNet-DQN)

PyTorch强化学习实战(14)------优先经验回放机制

PyTorch强化学习实战(15)------Dueling DQN

PyTorch强化学习实战(16)------Categorical DQN

PyTorch强化学习实战(17)------强化学习训练加速

PyTorch强化学习实战(18)------基于DQN处理股票交易问题

PyTorch强化学习实战(19)------策略梯度法

PyTorch强化学习实战(20)------优势演员-评论家(Advantage Actor-Critic, A2C)

PyTorch强化学习实战(21)------异步优势演员-评论家(Asynchronous Advantage Actor-Critic, A3C)

PyTorch强化学习实战(22)------将强化学习应用于TextWorld互动小说游戏

PyTorch强化学习实战(23)------强化学习在网页导航中的应用

相关推荐
Mr数据杨1 小时前
乌克兰新闻来源分类实战 从文本分类到媒体识别建模
人工智能·数据分析·kaggle竞赛
Mid_search1 小时前
Dueling Network
人工智能·深度学习·强化学习·dueling network
大唐荣华1 小时前
模型训练显卡选型深度指南:价格、数量、性能与自建租赁抉择
人工智能·机器人·模型训练·具身智能·avida·navida
万岳科技系统开发1 小时前
私域直播AI数字人系统:AI技术如何重构企业直播运营模式
大数据·人工智能·重构
AI_Cloud_推荐1 小时前
SpringBoot集成百度人脸识别SDK实战:人脸检测、比对与注册
大数据·人工智能·spring boot·安全·百度·视觉检测
u1301301 小时前
AI 日报(2026年9月3日)
人工智能
AI 思录1 小时前
“人眼看着正常,AI执行却出事“:Prompt事故档案(一)
人工智能·安全·prompt·用户体验·ai伦理
Rauser Mack1 小时前
从交互困境到语音闭环:桌面AI全语音数字秘书架构解析
人工智能·架构·交互