Hungry Hungry Hippos: Towards Language Modeling with State Space Models

发表时间: 2022-12 · arXiv:2212.14052

原文: https://arxiv.org/abs/2212.14052

速读

一句话结论 本文提出了新型状态空间模型层 H3 与高效训练算法 FlashConv,在解决传统状态空间模型表达能力不足与硬件利用率低的问题后,成功训练出 2.7B 参数的语言模型,在困惑度上超越同等规模的 Transformer,并实现了最高 2.4 倍的推理提速。

要解决什么问题 状态空间模型(SSM)具有 $O(N \log N)$ 的计算复杂度,理论上优于 Transformer,但在语言建模中面临两个卡点。一是表达能力不足。通过归纳头和关联回忆等合成任务发现,现有 SSM 缺乏两种关键能力:无法在特定事件后回忆早期词元,也无法跨序列比较词元,导致其困惑度显著落后于注意力机制。二是硬件效率障碍。标准快速傅里叶变换(FFT)卷积是 IO 密集型的,需频繁在 GPU 内存中读写中间结果,无法有效利用张量核心等专用矩阵乘法单元,导致实际训练速度慢于 Transformer。当序列超出 GPU 高速缓存容量时,内存瓶颈更会直接阻断模型扩展。

怎么做的 核心思路是通过结构化设计弥补状态空间模型的表达缺陷,并利用软硬件协同优化打破内存读写瓶颈。针对表达能力卡点,作者设计了 H3 层,其结构受线性注意力启发,由两个具有特定矩阵约束的离散状态空间模型与乘法交互机制构成。具体而言,输入序列首先被投影为 $Q$、$K$、$V$ 三个信号。第一个部件是移位状态空间模型,其状态矩阵被约束为移位矩阵,负责将状态向量逐位移动,充当检测特定事件并记录近期词元的短期记忆。第二个部件是对角状态空间模型,其状态矩阵被约束为对角矩阵,负责在整个序列中持久化地记住词元。为了实现跨序列的词元比较,H3 引入了输入投影与模型输出之间的乘法交互。当头维度为 1 时,H3 层的核心计算可定义为: $$O = \text{SSM}_{diag}(\text{SSM}_{shift}(K) \odot V) \odot Q$$ 其中 $\odot$ 表示逐点乘法。移位模型的输出与 $V$ 的乘法交互实现了局部比较并控制信息是否流入对角模型,而对角模型的输出与 $Q$ 的乘法交互则实现了全局序列的匹配提取。针对硬件效率,作者提出 FlashConv 算法。对 8K 以内的序列,采用融合的分块 FFT:先通过核函数融合将 FFT、逐点乘法和逆变换合并,在 GPU 高速缓存中完成以消除 IO 瓶颈;再利用 Cooley-Tukey 分解将变换拆解为块对角矩阵乘法,激活张量核心算力。对于超过 8K 的超长序列,算法引入状态传递机制。它将长序列切分为大小为 $N_0$ 的数据块,利用状态空间模型的循环特性,在处理完第 $c-1$ 个块后,将其最终状态 $x_{N_0}^{(c-1)}$ 传递给第 $c$ 个块作为初始条件。状态的跨块更新公式定义为: $$x_{N_0}^{(c)} = A^{N_0}x_{N_0}^{(c-1)} + M_{ux}u^{(c)}$$ 这种分块计算与状态传递的结合,使得模型在保持近线性复杂度的同时,能够无限扩展序列长度。

效果如何 实验在 A100 GPU 集群上进行,模型从 125M 扩展至 2.7B 参数,在 400B tokens 的 The Pile 数据集上训练。对比基线涵盖三条路线:代表标准自回归路线的 Transformer 模型(GPT-2、GPT-Neo、OPT),代表原有状态空间模型路线的 S4D 与 GSS,以及代表高效注意力路线的 Performer、Reformer 和线性注意力。量化结果显示,在语言建模任务上,仅包含两层注意力机制、其余全为 H3 层的混合 H3-注意力模型展现出极强竞争力。在 The Pile、OpenWebText 等数据集上,2.7B 参数的混合模型在困惑度上全面超越或持平于同等规模的 Transformer。在 SuperGLUE 的零样本和 3 样本分类任务中,混合模型在过半数任务上击败最优的 Transformer 基线。在硬件效率方面,得益于 FlashConv 加持,1.3B 参数的混合模型在生成 128 个词元的推理任务中,吞吐量最高达同等规模 Transformer 的 2.4 倍,且序列越长优势越显著;在 Long Range Arena 基准测试中,FlashConv 也让原有的 S4 模型提速 2 倍。局限性在于:在部分零样本生成任务中,模型容易生成不相关的长文本,需依赖少样本示例引导才能输出正确格式;此外,目前性能最强的仍是保留少量注意力层的混合架构,纯 H3 模型的表达能力仍有探索空间。

A1 主要贡献

本文旨在解决状态空间模型(SSM)在语言建模领域相较于Transformer存在的两个核心问题:模型表达能力不足和硬件利用率低下导致的训练速度慢。

核心问题与研究目标:

  1. 表达能力差距:尽管SSM在某些模态(如时间序列、音频)上表现出色,但在语言建模上,其性能(以困惑度PPL衡量)显著落后于Transformer。本文旨在探究此差距是否源于注意力机制固有的归纳偏置和能力。
  2. 硬件效率障碍:SSM的计算复杂度随序列长度呈近线性增长($O(N \log N)$),优于Transformer的二次方增长($O(N^2)$)。然而,由于未能有效利用现代硬件(如GPU的张量核心),SSM的实际运行速度反而更慢。本文致力于提升SSM在现代加速器上的训练效率。

主要创新与贡献:

  1. 通过合成任务揭示SSM的表达能力缺陷

    • 本文利用被认为是Transformer上下文学习能力基础的合成语言建模任务,来评估SSM与注意力机制的差距。
    • 研究发现,现有SSM在两项关键能力上存在不足:回忆序列中早期的词元(token)跨序列比较词元
  2. 提出新型SSM层——H3 (Hungry Hungry Hippo)

    • 为解决上述表达能力缺陷,本文设计了H3层。该层通过堆叠两个具有特定结构的SSM(一个使用移位矩阵,一个使用对角矩阵),并引入输入投影与其输出之间的乘法交互,来显式地增强模型的回忆和比较能力。
    • 效果:H3在合成语言任务上表现与注意力机制相当。在OpenWebText数据集上,纯H3模型与Transformer的困惑度差距从3.4缩小到0.4。更重要的是,一个包含两个注意力层的混合H3-注意力模型在OpenWebText上的表现甚至比纯Transformer模型还要好1.0 PPL。

    图1:左:H3堆叠了两个带有移位和对角矩阵的离散SSM,并使用输入投影及其输出之间的乘法交互来模拟序列中点与点之间的比较。中:H3可以执行关联回忆——这对注意力来说很容易,但对现有的SSM来说则不然。右:FlashConv使用一种新的状态传递算法,通过融合的分块FFTConv来提高SSM的硬件效率,使H3能够扩展到十亿参数级别的模型。
    图1:左:H3堆叠了两个带有移位和对角矩阵的离散SSM,并使用输入投影及其输出之间的乘法交互来模拟序列中点与点之间的比较。中:H3可以执行关联回忆——这对注意力来说很容易,但对现有的SSM来说则不然。右:FlashConv使用一种新的状态传递算法,通过融合的分块FFTConv来提高SSM的硬件效率,使H3能够扩展到十亿参数级别的模型。

  3. 提出高效训练算法——FlashConv

    • 为提升SSM的硬件效率,本文提出了FlashConv算法,其灵感来源于IO感知的注意力机制【15, Flashattention: Fast and memory-efficient exact attention with io-awareness, 2022, Advances in Neural Information Processing Systems】。
    • 核心技术
      • 融合的分块FFT(Fused Block FFT):对于中等长度序列(在A100上可达8K),该算法将FFT计算分解为一系列矩阵乘法,从而利用GPU的专用矩阵乘法单元(如Tensor Cores),并通过核函数融合(kernel fusion)减少内存读写开销。
      • 状态传递算法(State Passing Algorithm):对于超长序列(>8K),该算法利用SSM的循环特性,将输入序列分块处理。每处理完一个数据块,就计算并传递一个状态向量给下一个数据块,从而在保持近线性计算复杂度的同时,将模型扩展到任意序列长度。
    • 效果:FlashConv在Long Range Arena基准测试上取得了2倍的速度提升,并使混合语言模型的文本生成速度比Transformer快2.4倍。
  4. 大规模语言模型验证

    • 利用FlashConv,本文成功将混合H3-注意力语言模型扩展至2.7B参数,并在The Pile数据集上进行了训练。
    • 结果:这些模型在困惑度上优于同等规模的Transformer,并在SuperGLUE基准测试的大多数任务中,在零样本和少样本学习方面与Transformer相当或更优。

A3 背景知识

本节介绍状态空间模型(SSM)和线性注意力的背景,它们是H3层的设计灵感来源。

2.1 状态空间模型

连续时间状态空间表示

离散时间状态空间表示

SSM作为卷积

通过FFT实现SSM

2.2 线性注意力

线性注意力与RNN的联系

A2 方法细节

3 Hungry Hungry Hippos层:用于建模离散序列

为了理解SSM和注意力在语言建模上的差距,我们研究了两个合成语言建模任务。这些任务启发我们设计了H3层,通过增加一个基于移位矩阵的离散SSM和乘法交互来有效地建模离散序列。我们接着展示了H3层具有足够的表达能力来解决这些合成任务,并且这种理解也带来了在真实语言建模基准上更好的性能。

3.1 动机:合成语言建模任务

合成任务描述

表1:合成语言建模任务。

现有SSM的局限性

表2:两层模型在合成语言任务上的评估。

3.2 H3层

H3层设计

高层直觉

H3层与线性注意力的联系

移位和对角SSM:记忆关键Token

乘法交互:实现比较功能

H3层算法

算法1 H3层

效率分析

3.3 表达能力

H3的表达能力验证

H3解决关联回忆任务的机制

从合成语言到自然语言的转化

表3:SSM变体与Transformers在OpenWebText上的困惑度对比。所有模型均为12层,大小约为125M,并使用相同的超参数训练50B个词元。

扩展:H3-注意力混合模型

4 FlashConv:高效训练SSM

FlashConv简介

4.1 融合的分块FFTConv

短序列加速技术

核函数融合

分块FFT

4.2 状态传递

长序列处理的挑战与解决方案

状态传递算法细节

算法2 状态传递算法

算法正确性

A4 实验

实验环境

  1. 数据集

    • 语言建模
      • The Pile【21, The pile: An 800gb dataset of diverse text for language modeling, 2021, arXiv preprint】: 用于训练125M至2.7B参数模型的核心数据集,训练量达400B tokens。
      • OpenWebText【23, Openwebtext corpus, 2019】: 用于125M模型与Transformer的详细对比。
      • WikiText-103【43, Pointer sentinel mixture models, 2016】: 用于评估模型的零样本迁移能力。
      • PG-19【54, Compressive transformers for long-range sequence modelling, 2019, ICLR】: 用于评估长文本建模能力。
    • 合成任务
      • Induction Head & Associative Recall: 用于诊断模型(特别是H3)的表达能力。
    • 长序列基准
      • Long Range Arena (LRA)【59, Long range arena: A benchmark for efficient transformers, 2020, ICLR】: 用于评估FlashConv的加速效果。
    • 非文本序列建模
      • TUSZ v1.5.2 EEG Corpus【56, The temple university hospital seizure detection corpus, 2018, Frontiers in neuroinformatics】: 用于癫痫分类任务。
      • Speech Commands (SC10)【64, Speech commands: A dataset for limited-vocabulary speech recognition, 2018, arXiv preprint】: 用于原始音频分类任务。
      • fMRI数据集【60, Self-supervised learning of brain dynamics from broad neuroimaging data, 2022, arXiv preprint; 61, The wu-minn human connectome project: an overview, 2013, Neuroimage; 38, Functional boundaries in the human cerebellum revealed by a multi-domain task battery, 2019, Nature neuroscience】: 用于脑功能状态解码。
  2. 模型架构

    • H3-Attention Hybrid Models: 在125M, 355M, 1.3B, 2.7B四个尺寸上进行训练。
      • 125M: 12层,隐层维度1024,MLP维度4096,12个头。注意力层在第1和第7层。
      • 355M: 24层,隐层维度1024,MLP维度4096,16个头。注意力层在第1和第13层。
      • 1.3B: 24层,隐层维度2048,MLP维度8192,16个头。注意力层在第1和第13层。
      • 2.7B: 32层,隐层维度2560,MLP维度10240,20个头。注意力层在第10和第21层。
    • H3特定参数:SSM状态大小为64。混合模型中H3的头维度为1,纯H3模型为8。
    • 基线模型: GPT-2, GPT-Neo, OPT等同等规模的Transformer模型。
  3. 硬件配置

    • GPU: 训练在单个包含16块A100-40GB GPU的节点或一个由8块A100-80GB GPU组成的集群上进行。
    • 基准测试: 速度测试在A100-SMX4-40GB GPU上进行。
  4. 软件配置

    • 代码实现: 基于PyTorch,使用混合精度训练(AMP),其中MLP和注意力部分使用bf16,FFTConv部分使用fp32。
    • Tokenizer: 使用GPT-2 BPE tokenizer。
    • 优化器: AdamW。
    • 依赖库: HuggingFace Transformers【65, Transformers: State-of-the-art natural language processing, 2020, EMNLP】。

实验结果

H3 语言建模评估 (Section 5)

本节评估混合H3-注意力模型在困惑度、零样本/少样本学习以及推理速度方面与Transformer的对比。

  1. 困惑度 (Perplexity)

    • 实验内容:在The Pile、OpenWebText和WikiText-103上评估了125M到2.7B参数的混合H3模型,并与GPT-Neo和GPT-2进行比较。
    • 实验结果
      • 在The Pile上,125M的混合H3模型优于同样在该数据集上训练的GPT-Neo(表4)。
      • 在向OpenWebText和WikiText-103的零样本迁移任务中,混合H3模型同样优于GPT-Neo和GPT-2(表4)。
    • 分析结论:混合H3模型在困惑度指标上全面超越或持平于同等规模的Transformer模型,证明了H3层在提升语言建模能力上的有效性。

    表4:模型在The Pile, OpenWebText和WikiText-103上的困惑度(越低越好)。GPT-Neo和混合H3模型在The Pile上训练,而GPT2在WebText上训练。所有模型使用相同的GPT2分词器。我们报告了GPT-2模型在The Pile上的困惑度(*)作为参考,但由于训练数据不同,性能不直接可比。

  2. 零样本和少样本学习性能

    • 实验内容:在SuperGLUE基准上,比较了混合H3模型与OPT、GPT-Neo、GPT-2的零样本和3样本性能。
    • 实验结果
      • 零样本:混合H3模型在超过一半的任务上表现优于或持平于最好的Transformer基线(表5)。
      • 3样本:趋势与零样本一致,混合模型在多数任务上取得领先或有竞争力的结果(表6)。
    • 分析结论:在下游任务的零/少样本学习上,混合H3模型展现出与Transformer相当甚至更强的能力,表明其学到的语言表示是有效且通用的。

    表5:在SuperGLUE上的零样本准确率(使用logit评分)。最好结果加粗,次好结果下划线。

    表6:在SuperGLUE上的3样本准确率(使用logit评分)。最好结果加粗,次好结果下划线。

  3. 推理速度

    • 实验内容:比较了1.3B参数的混合H3模型和Transformer在文本生成任务上的推理吞吐量。
    • 实验结果:由于SSM的循环特性,混合H3模型的推理吞吐量最高可达Transformer的2.4倍,且序列越长,优势越明显(表7)。
    • 分析结论:H3-Attention混合模型在保持高质量的同时,显著提升了生成任务的效率。

    表7:在A100 80GB上,1.3B模型的推理吞吐量。批量大小为64,提示长度为512、1024或1536,每个序列生成128个token。混合H3的推理速度比同等规模的Transformer快高达2.4倍。序列越长,差异越大。

FlashConv 效率评估 (Section 6)

本节评估FlashConv对SSM的加速效果。

  1. Long Range Arena (LRA) 基准测试

    • 实验内容:使用FlashConv加速S4模型,并在LRA基准上测试其性能。
    • 实验结果:FlashConv将S4的速度提升了2倍,总体性能比Transformer快5.8倍,创造了新的SOTA速度记录(表8)。
    • 分析结论:FlashConv能有效加速现有的SSM模型,在长序列任务上展现出巨大优势。

    表8:在LRA基准测试上的加速效果。

  2. H3模块与Attention的速度对比

    • 实验内容:在不同序列长度(256到32K)下,测试了使用FlashConv的H3模块前向和后向传播的时间,并与FlashAttention进行对比。
    • 实验结果
      • FlashConv(结合分块FFT和状态传递)比基于cuFFT的朴素FFTConv实现快2-3倍(图2)。
      • 短序列 (<=512): 内核融合提供了高达3.4倍的加速。
      • 中序列 (1k-8k): 分块FFT提供了高达2倍的加速。
      • 长序列 (>=16k): 状态传递算法使得FFTConv快了2.3倍。
      • 随着序列长度增加,H3的运行时间呈近线性增长,而Attention则呈二次方增长,使得H3在长序列上比最快的Attention实现快数十倍(图2)。
    • 分析结论:FlashConv通过多层次优化(内核融合、分块FFT、状态传递)显著提升了SSM的硬件效率,使其在各种序列长度上都比Attention更快,尤其是在长序列场景下。

    图2:我们比较了不同算法执行基于FFT的卷积的速度,以及FlashAttention【15, Flashattention: Fast and memory-efficient exact attention with io-awareness, 2022, Advances in Neural Information Processing Systems】(我们所知的最快的注意力实现)。我们使用批量大小8,隐藏维度1024,在A100-SMX4-40GB GPU上测量了从256到32k的不同序列长度。我们看到,对于短序列(最大512),内核融合比朴素的FFTConv快达3.4倍;对于中等长度序列(1k-8k),分块FFT快达2倍;对于长序列(16k及以上),状态传递使得FFTConv快了2.3倍。
    图2:我们比较了不同算法执行基于FFT的卷积的速度,以及FlashAttention【15, Flashattention: Fast and memory-efficient exact attention with io-awareness, 2022, Advances in Neural Information Processing Systems】(我们所知的最快的注意力实现)。我们使用批量大小8,隐藏维度1024,在A100-SMX4-40GB GPU上测量了从256到32k的不同序列长度。我们看到,对于短序列(最大512),内核融合比朴素的FFTConv快达3.4倍;对于中等长度序列(1k-8k),分块FFT快达2倍;对于长序列(16k及以上),状态传递使得FFTConv快了2.3倍。

附加实验 (Appendix F)

A5 结论

本文旨在理解并缩小注意力机制与状态空间模型(SSM)在语言建模领域的模型能力和硬件效率差距。

主要研究成果:

  1. 模型设计 (H3):通过合成语言任务的探索,我们发现现有SSM在回忆比较能力上存在不足。为此,我们设计了H3层,该层通过堆叠两个SSM并引入乘法交互,显著提升了SSM的表达能力,使其在语言建模任务上表现出与注意力机制惊人地具有竞争力。
  2. 算法优化 (FlashConv):我们提出了BlockFFTConv算法,该算法利用GPU的矩阵乘法单元以及SSM的循环-卷积双重视图,大幅提升了SSM的计算速度,从而降低了注意力与SSM之间的硬件壁垒。

未来工作展望:

A6 附录

B 线性注意力与时变系统

将线性注意力与LTI系统和SSM建立联系


这与SSM(第2节)的形式相同,只是矩阵可以依赖于时间步。


其中$A, B, C$是学习的矩阵。然后我们通过将$y_{i+1}$与$\phi(Q_i)^T$相乘来进行后处理。

C 方法细节

C.1 反向传播


由于当$i \ge L$时$u'[i] = f'[i] = 0$,我们可以填入矩阵中的零:


其中我们用$u'[-i]$表示$u'[2L-i]$。注意到这个矩阵与$H_{u'}$具有相同的格式!令$u'^* = [u'[0], u'[-1], ..., u'[-(2N-1)]]$。那么:


其中$^*$代表复共轭。我们可以利用这个性质高效地计算$df'$:

其中$FFT^*$表示取FFT的复共轭,$dy'$表示用零填充的$dy$。

C.2 状态传递矩阵

D 证明

D.1 H3 表达能力

D.2 注意力表达能力

D.3 H3 复杂度

D.4 状态传递正确性

F 附加实验

F.1 LRA 准确率

表9:H3与S4D在LRA上的性能对比。

F.2 WikiText103

表10:WikiText103上的测试PPL。

F.3 PG-19

表11:PG-19上的测试PPL。

F.4 长度外推

表12:在长度为20的序列上训练的H3模型,在长度为20和40的序列上评估的关联回忆准确率。

F.5 按Token数量扩展

表13:在The Pile上用较少token训练的模型的测试PPL。

F.6 H3语言模型

表14:SuperGLUE上使用排名分类的零样本性能。每个模型大小的最佳结果加粗。


表15:SuperGLUE上使用排名分类的3样本性能。每个大小的最佳结果加粗,次佳结果下划线。

F.7 生成性能

表16:SuperGLUE上的零样本性能。每个大小的最佳结果加粗,次佳结果下划线。


表17:SuperGLUE上使用生成的3样本性能。每个大小的最佳结果加粗,次佳结果下划线。

F.8 非文本序列建模

表18:从原始EEG(序列长度12000)进行60秒癫痫分类的性能(AUROC)。

表19:在原始音频(序列长度16000)上的SC 10类分类。

图3:模型训练过程中,训练和评估数据集的上游平均绝对误差(Lrec)。
图3:模型训练过程中,训练和评估数据集的上游平均绝对误差(Lrec)。

图4:最终预训练模型在每个脑区体素的平均绝对误差(Lrec),投射到FsAverage模板的膨胀皮质表面。
图4:最终预训练模型在每个脑区体素的平均绝对误差(Lrec),投射到FsAverage模板的膨胀皮质表面。
**表20:在fMRI数据上预训练的模型的下游适应性能,对20次不同随机种子的训练运行取平均。F1分数是宏平均。**