序列模型与Transformer

闪电注意力(flash attention)

/ flash uh-TEN-shun /

闪电注意力,是一种巧妙的办法,它通过聪明地安排数据在芯片内部的搬运,算出与原来完全相同的注意力结果,却快得多、省内存得多。关键的领悟是:在现代 GPU 上,慢的那部分通常不是算术——而是把数据在芯片那一小块快如闪电的片上内存、与它更大却慢得多的主内存之间来回倒腾。标准注意力会把一个巨大的分数矩阵写进慢内存再读回来,而这趟流量才是真正的瓶颈。

闪电注意力压根不在慢内存里搭出那个完整矩阵。它把序列分成一块块能塞进快速片上内存的小瓦片,为每块算注意力,再用一个边算边更新、数值上很谨慎的 softmax 版本,巧妙地把部分结果即时合并起来。它的输出在数学上与普通注意力一模一样——并非近似——但因为它从不实体化那个巨大的中间矩阵,所用内存只随序列长度线性增长,而非平方增长,实践中还快上好几倍。

这是一个「在丝毫不改数学的前提下,改变了可能性边界」的教科书例子。通过砍掉内存代价,闪电注意力使得在远比从前更长的序列上训练和运行模型变得可行,直接成全了今天人们用的许多长上下文模型。诚实的说法是:它并没有打破注意力那根本的平方计算量——运算的次数还在那里——它移除的是内存墙与慢内存流量,而那才是真正的实践上限。这是一场系统的胜利,而非关于注意力本身的新算法。

在一段 16000 词元的序列上,标准注意力试图在内存里端着一个 16000×16000 的分数矩阵——数以亿计的数字——于是憋住了。闪电注意力压根不搭它;它流式地穿过一块块小瓦片,产出一模一样的输出,却能舒舒服服地装进快速片上内存。

同样的答案,没有那个巨大矩阵——一个内存技巧,而非新公式。

一个常见的夸大,是说闪电注意力「解决」了注意力的平方代价。它并没有——算术运算的次数没变。它移除的是内存瓶颈与慢内存流量,那才是绑手绑脚的实践上限。它算出的是完全相同的结果,只是效率高得多。

又称
FlashAttention闪电注意力閃電注意力IO-aware attention