AGaLiTe: Approximate Gated Linear Transformers for Online Reinforcement Learning

📄 arXiv: 2310.15719v2 📥 PDF

作者: Subhojeet Pramanik, Esraa Elelimy, Marlos C. Machado, Adam White

分类: cs.LG, cs.AI

发布日期: 2023-10-24 (更新: 2024-10-15)

备注: Published in Transactions on Machine Learning Research


💡 一句话要点

提出AGaLiTe以解决在线强化学习中的变压器架构问题

🎯 匹配领域: 支柱二:RL算法与架构 (RL & Architecture)

关键词: 在线强化学习 变压器架构 自注意力机制 递归神经网络 长距离依赖性 性能提升 内存优化

📋 核心要点

  1. 现有变压器架构在在线强化学习中存在自注意力机制需访问全部历史信息和推理成本高的问题。
  2. 本文提出AGaLiTe,通过递归替代自注意力机制,实现上下文无关的推理成本,提升在线强化学习性能。
  3. 实验结果表明,AGaLiTe在推理成本上至少降低40%,内存使用减少超过50%,在更困难任务中性能提升超过37%。

📝 摘要(中文)

本文研究了针对部分可观察在线强化学习设计的变压器架构。变压器的自注意力机制能够捕捉长距离依赖性,这是其在处理序列数据时有效性的主要原因。然而,变压器在在线强化学习中的应用受到两个显著缺陷的限制:一是自注意力机制需要访问整个历史信息以提供上下文,二是推理成本昂贵。为此,本文提出了递归替代方案,提供上下文无关的推理成本,有效利用长距离依赖性,并在在线强化学习任务中表现良好。我们在诊断环境中量化了架构不同组件的影响,并在2D和3D像素基础的部分可观察环境中评估性能提升。与最先进的架构GTrXL相比,我们的方法推理成本至少降低40%,内存使用减少超过50%。在更困难的任务中,我们的方法在性能上提升超过37%。

🔬 方法详解

问题定义:本文旨在解决在线强化学习中变压器架构的两个主要问题:自注意力机制需要访问全部历史信息以提供上下文,以及推理成本过高,这限制了其在实际应用中的有效性。

核心思路:论文提出了一种递归替代方案,旨在降低推理成本并有效利用长距离依赖性。通过这种设计,AGaLiTe能够在不依赖完整历史信息的情况下进行推理,从而提高效率。

技术框架:AGaLiTe的整体架构包括多个模块,主要包括递归自注意力机制、状态表示模块和决策模块。该架构能够在部分可观察环境中进行有效的学习和推理。

关键创新:AGaLiTe的核心创新在于引入递归机制替代传统的自注意力机制,显著降低了推理成本,并减少了对历史信息的依赖。这一设计使得模型在处理长序列时更加高效。

关键设计:在参数设置上,AGaLiTe优化了递归层的数量和每层的神经元数量,同时采用了适应性损失函数以提高训练效率。网络结构上,AGaLiTe结合了长短期记忆(LSTM)和卷积神经网络(CNN)以增强特征提取能力。

🖼️ 关键图片

fig_0
fig_1
fig_2

📊 实验亮点

实验结果显示,AGaLiTe在推理成本上至少降低40%,内存使用减少超过50%。在更复杂的任务中,AGaLiTe的性能提升超过37%,相较于现有的GTrXL架构表现出色,证明了其在在线强化学习中的有效性。

🎯 应用场景

AGaLiTe的研究成果在多个领域具有潜在应用价值,包括机器人控制、游戏智能体、自动驾驶等。通过提高在线强化学习的效率,该方法能够在实时决策和动态环境中实现更优的性能,推动智能系统的进一步发展。

📄 摘要(原文)

In this paper we investigate transformer architectures designed for partially observable online reinforcement learning. The self-attention mechanism in the transformer architecture is capable of capturing long-range dependencies and it is the main reason behind its effectiveness in processing sequential data. Nevertheless, despite their success, transformers have two significant drawbacks that still limit their applicability in online reinforcement learning: (1) in order to remember all past information, the self-attention mechanism requires access to the whole history to be provided as context. (2) The inference cost in transformers is expensive. In this paper, we introduce recurrent alternatives to the transformer self-attention mechanism that offer context-independent inference cost, leverage long-range dependencies effectively, and performs well in online reinforcement learning task. We quantify the impact of the different components of our architecture in a diagnostic environment and assess performance gains in 2D and 3D pixel-based partially-observable environments (e.g. T-Maze, Mystery Path, Craftax, and Memory Maze). Compared with a state-of-the-art architecture, GTrXL, inference in our approach is at least 40% cheaper while reducing memory use more than 50%. Our approach either performs similarly or better than GTrXL, improving more than 37% upon GTrXL performance in harder tasks.