直接回答:标准注意力的瓶颈不是算力而是显存读写:它要把 N×N 的注意力分数矩阵完整地写入 HBM、做 softmax、再读回来乘 V,序列一长就是典型的访存受限(memory-bound)操作,算力大量闲置。FlashAttention(Dao et al., 2022)的核心是 IO 感知的精确注意力:利用 GPU 显存层级中 SRAM 快而小的特点,把 Q、K、V 分块(tiling)搬进片上 SRAM,在块内配合在线 softmax(online softmax)增量地计算并归一化注意力输出,全程不把完整的 N×N 矩阵写回 HBM。结果是显存占用从 O(N²) 降到 O(N),实测提速 2–4 倍,且数值结果与标准注意力完全等价——它不是近似算法。

展开解析:在线 softmax 的关键技巧是维护运行中的行最大值 m 和归一化分母 l:每处理一个新的 K/V 块,用新旧最大值之差对已有累积输出做指数缩放再合并,单次遍历即可完成 softmax 且数值稳定,避免了对整行分数的两次读写。反向传播不保存注意力矩阵,只保存输出和 (m, l),反向时重算——用少量计算换显存,与梯度检查点思想相通。FlashAttention-2 进一步把并行维度扩展到序列方向、减少非 matmul 的低效 FLOPs 占比;FlashAttention-3 针对 Hopper 架构利用 TMA 异步拷贝、warp 专门化和 FP8 再进一步提速。它已成为所有主流训练与推理框架的默认注意力实现。

# 在线 softmax 的块合并逻辑(伪代码)
for kj, vj in blocks(K, V):
    s = q @ kj.T
    m_new = max(m, rowmax(s))
    p = exp(s - m_new)
    l = l * exp(m - m_new) + rowsum(p)
    o = o * exp(m - m_new) + p @ vj
    m = m_new
o = o / l

追问方向:为什么反向传播用重算而不是存矩阵?分块大小如何依据 SRAM 容量确定?FlashAttention 与稀疏注意力、滑窗注意力能否结合?(约 650 字)