模型读一段长文本时,注意力层要计算各个 Token 之间的关系。序列越长,中间结果越大,GPU 花在搬运数据上的时间也会增加。FlashAttention 盯住的正是这些显存读写。它仍然计算精确注意力,没有靠删掉大部分关系来换速度。
我读了原论文、FlashAttention-2 论文和项目官方实现。原论文把 GPU 内存分成速度不同的层级。容量较大的高带宽显存适合存数据,片上 SRAM 更快却很小。普通实现会把较大的注意力矩阵写回显存,再读出来继续算。FlashAttention 将计算分块,让一小块数据留在更快的片上内存中完成更多步骤,减少往返搬运。
它改的是计算过程
这件事容易和模型压缩混在一起。模型量化会降低权重或计算使用的数值精度,FlashAttention 改写注意力算子的执行顺序。模型参数没有因此减少,注意力结果也仍与标准实现一致,只允许正常的浮点误差。两项优化可以同时使用,验收时应分别确认精度、显存和速度。
长序列更容易看到显存收益。官方仓库给出的 A100 基准中,在特定批次、头维度和数据类型下,序列长度 2K 时中间显存约省十倍,4K 时约省二十倍。这个数字展示了增长趋势,不能直接拿去估算另一块显卡。批大小、是否训练、是否带 dropout,以及框架选到哪个内核都会改变结果。
FlashAttention-2 又改了 GPU 上的工作分配。论文减少了不属于矩阵乘法的运算,并让线程块和 warp 更合理地分担一个注意力头。论文环境里,它相对第一版约快两倍,在 A100 上达到理论浮点吞吐的五成到七成多。这里测的是指定模型和硬件,业务里的端到端时间还包含分词、数据传输、其他网络层和输出采样。
训练和在线生成要分开测
训练会保存反向传播需要的中间状态,长序列带来的显存压力很明显。在线生成每次通常只增加一个 Token,主要工作会转向读取已有的 KV cache。官方实现为这种迭代解码提供了专门接口,也支持更新 KV cache。一个训练基准很快,不能证明同一配置的在线首字时间和每 Token 延迟都同样改善。
硬件支持也有边界。官方仓库分别列出 CUDA 与 ROCm 后端、支持的 GPU、数据类型和 head dimension。某些功能只在较新的架构或版本上可用,Windows 编译仍标着需要更多测试。框架还可能根据输入自动选用别的注意力实现,安装了包也未必每次都走 FlashAttention。
实际接入时,我会先固定模型、输入长度和批大小,记录普通注意力的峰值显存与端到端时间,再启用目标内核重复测试。随后核对日志或性能分析结果,确认框架确实选到了预期实现。显存降下来以后,再尝试加长上下文或扩大批次。这样得到的收益属于自己的机器和请求,不会把论文数字当成采购承诺。