长期信息1 分钟阅读·
IO感知切块:注意力加速的新原则
Transformers在长序列上慢且耗内存,因为自注意力的时间和内存复杂度随序列长度二次增长。
证据E3 可检验写法标准
Transformers在长序列上慢且耗内存,因为自注意力的时间和内存复杂度随序列长度二次增长。近似注意力方法试图用模型质量换计算量,但往往没有带来墙钟时间上的加速。
论文提出一个缺失原则:让注意力算法感知IO,即考虑GPU不同层级内存之间的读写。FlashAttention用切块(tiling)减少GPU高带宽内存(HBM)与片上SRAM之间的读写次数。
它声称是精确注意力算法,不需要近似,结果与标准算法一致。同时扩展出近似版本,比所有现有近似方法更快。
在GPT-2(序列长度1K)上3倍加速,在long-range arena(序列长度1K-4K)上2.4倍加速。GPT-2困惑度改善0.7,长文档分类提升6.4个百分点。首次使Transformer在Path-X(16K序列,61.4%)和Path-256(64K序列,63.1%)上超过随机水平。
实验主要在GPU上、特定模型和序列长度上进行。超过一定长度后部分近似方法可能反超。不能据此断言近似方法普遍无效,也不能外推到非GPU硬件或未测试的配置。
假设你正在训练一个长上下文模型,发现GPU利用率不高但内存带宽打满——这正是IO瓶颈的典型表现。此时减少HBM与SRAM之间的搬运比增加算力更有效。
《FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness》,Tri Dao等,2022年NeurIPS。适合想理解为什么长序列Transformer慢、怎么在不牺牲精度的情况下加速的读者。从摘要和IO分析部分入手。