Simple Hardware-Efficient Long Convolutions for Sequence Modeling

Simple Hardware-Efficient Long Convolutions for Sequence Modeling
复制标题

DOI:
10.48550/arxiv.2302.06646
复制
发表时间:
2023-02
期刊:
ArXiv
影响因子:
--
通讯作者:
Daniel Y. Fu;Elliot L. Epstein;Eric N. D. Nguyen;A. Thomas;Michael Zhang;Tri Dao;A. Rudra;Christopher Ré
Daniel Y. Fu;Elliot L. Epstein;Eric N. D. Nguyen;A. Thomas;Michael Zhang;Tri Dao;A. Rudra;Christopher Ré
中科院分区:
其他
文献类型:
--
作者:
Daniel Y. Fu;Elliot L. Epstein;Eric N. D. Nguyen;A. Thomas;Michael Zhang;Tri Dao;A. Rudra;Christopher Ré

文献摘要

被引文献

相似文献

状态空间模型 (SSM) 在长序列建模方面具有高性能,但需要复杂的初始化技术和专门的实现来实现高质量和运行时性能。我们研究了一种简单的替代方案是否可以在性能和效率上与 SSM 相匹配:直接学习序列上的长卷积。我们发现实现高性能的关键要求是保持卷积核的平滑。我们发现简单的干预措施(例如压缩内核权重)可以使内核变得平滑,并恢复 SSM 在一系列任务上的性能,包括远程领域、图像分类、语言建模和大脑数据建模。接下来,我们开发了 FlashButterfly,这是一种 IO 感知算法,用于提高长卷积的运行时性能。 FlashButterfly 诉诸卷积的经典 Butterfly 分解,以减少 GPU 内存 IO 并提高 FLOP 利用率。 FlashButterfly 将卷积速度提高了 2.2$\times$,并允许我们在 Path256 上进行训练,这是一项具有挑战性的序列长度 64K 的任务,我们将最先进的技术提高了 29.1 个点,同时训练速度比之前的工作快了 7.2$\times$。最后,我们引入了 FlashButterfly 的扩展,它可以学习 Butterfly 分解的系数,从而在不增加运行时间的情况下提高表现力。使用此扩展,我们在 WikiText103 上的性能优于 Transformer 0.2 PPL,参数减少了 30%。
State space models (SSMs) have high performance on long sequence modeling but require sophisticated initialization techniques and specialized implementations for high quality and runtime performance. We study whether a simple alternative can match SSMs in performance and efficiency: directly learning long convolutions over the sequence. We find that a key requirement to achieving high performance is keeping the convolution kernels smooth. We find that simple interventions--such as squashing the kernel weights--result in smooth kernels and recover SSM performance on a range of tasks including the long range arena, image classification, language modeling, and brain data modeling. Next, we develop FlashButterfly, an IO-aware algorithm to improve the runtime performance of long convolutions. FlashButterfly appeals to classic Butterfly decompositions of the convolution to reduce GPU memory IO and increase FLOP utilization. FlashButterfly speeds up convolutions by 2.2$\times$, and allows us to train on Path256, a challenging task with sequence length 64K, where we set state-of-the-art by 29.1 points while training 7.2$\times$ faster than prior work. Lastly, we introduce an extension to FlashButterfly that learns the coefficients of the Butterfly decomposition, increasing expressivity without increasing runtime. Using this extension, we outperform a Transformer on WikiText103 by 0.2 PPL with 30% fewer parameters.