论文

VFA:通过全局最大预计算减轻 Flash Attention 中的向量运算

VFA: Relieving Vector Operations in Flash Attention with Global Maximum Pre-computation

模型推理推理加速

摘要

FlashAttention 式在线 softmax 通过片上内存流式传输分数图块并维护运行最大值和标准化器,从而可以使用线性内存进行精确的注意力计算。然而,随着注意力内核在现代加速器上接近峰值张量核心/立方体核心吞吐量,在线 softmax 的非 matmul 组件(尤其是每个图块的 rowmax 和 rowsum 缩减以及重新缩放链)可能会受到矢量或 SIMD 限制并主导延迟。本文重新审视了 FlashAttention,并提出了 Vector Relieved Flash Attention (VFA),这是一种硬件友好的方法,可以减少 rowmax 驱动的运行最大值更新,同时保留 online-softmax 结构。 VFA 通过关键块表示的廉价近似来初始化运行最大值,重新排序关键块遍历以优先考虑高影响的接收器和本地块,并冻结剩余块的最大值以避免重复减少和重新缩放。我们进一步将 VFA 与块稀疏跳跃方法(例如 BLASST)集成,形成向量缓解稀疏注意力(VSA),从而减少块计数和每个块的开销。值得注意的是,VFA和VSA完全避免了FA4.0中使用的更新阶段的条件重缩放操作。对 MMLU 和 MATH500 等基准的广泛评估以及注意力统计数据验证了我们的设计:(i)sink 和本地重新排序尽早稳定了运行最大值; (ii) 由于块内异质性,简单的 Q 和 K 块摘要失败; (iii) 当最大值出现在中间块时需要m初始化。总体而言,VFA 和 VSA 有效缓解了 online-softmax 缩减瓶颈,且没有性能损失。与 C16V32 基准相比,C8V32、C4V32 和 C4V16 在现代硬件上实现了近两倍的加速,同时遇到了矢量瓶颈。随着即将到来的架构改进,C4V16 将通过增强指数容量来提供六倍的加速。