
FlashAttention加速滑动窗口注意力prefill核心可以拆成三句话分块计算避免完整注意力矩阵落显存块级跳过避免计算窗口之外的Key块在线softmax保证跳过块之后数值依然正确。这三件事缺一不可尤其块级跳过是把滑动窗口的稀疏性真正变成收益的关键。这篇文章把这个组合的机制拆开讲清楚。你会理解prefill阶段为什么是长文本推理的瓶颈滑动窗口注意力在普通实现里为什么仍然慢以及FlashAttention加上窗口控制之后快在哪里、怎么验证、有哪些坑。适合正在做长文本推理、训练长序列模型或者想手写注意力kernel的开发者。1. 先把问题定义清楚prefill、滑动窗口和缓慢的真相1.1 Prefill为什么是长文本推理的第一道坎大模型推理通常分成两个阶段。第一个阶段叫prefill拿到整段prompt把每个token的隐藏状态一次性算完同时生成KV cache供后续使用。第二个阶段叫decode从一个生成位置开始逐个token地输出每一步只处理当前token。prefill最大的特点是并行度高所有prompt token同时参与计算。但正因为并行attention部分特别贵。对一个长度为n的prompt标准attention要计算n×n个注意力分数每个query位置都要和所有key位置做内积。n增长一倍计算量涨四倍。具体感受一下1万token的prompt注意力分数矩阵是1万×1万fp16下大概200MB。5万token就是5GB。10万token就超过20GB。这个矩阵还没算Q/K/V本身和KV cache单卡已经很难处理了。所以长文本场景里prefill往往是比decode更早撞到资源天花板的地方。1.2 滑动窗口注意力理论复杂度降了实际不一定滑动窗口注意力做的事情很简单每个token只和它前面连续W个token做attention窗口之外的Key位置完全不参与计算。这样理论计算量从O(n²)变成O(nW)序列越长、窗口越小收益越大。这个设计在很多长文本模型里已经用了。Mistral 7B的attention就是滑动窗口窗口大小4096。Longformer、BigBird这类稀疏注意力模型也大量使用局部窗口。从建模角度看它符合一个直觉文本依赖往往集中在局部不是每个token都需要看完整篇上下文。但这里要泼一盆冷水。理论复杂度降低了不代表实际运行更快。很多实现只是在PyTorch里写一个常规attention然后加一个mask把窗口外的位置变成负无穷。这种实现下n×n的分数矩阵照样会算出来照样会分配显存mask只是把不想要的位置遮住计算量一点都没少。我见过一个5万token的prompt配4096窗口的测试暴力mask方式下显存占用远超预期。原因很简单GPU在算QK^T的时候不管后面mask不mask整个矩阵都已经算完了。1.3 为什么需要kernel级稀疏而不是mask所以滑动窗口注意力要真正加速必须在kernel层面利用稀疏性窗口外的分数根本不算。这正是FlashAttention加窗口控制能解决的问题。FlashAttention本身并不改变attention的语义它只是改变计算和内存访问的方式。而滑动窗口是一个天然的稀疏模式两者结合的关键是让kernel在遍历Key块的时候只加载和计算窗口内相关的块。这个思路说起来简单落地时却涉及块区间计算、边界掩码、在线softmax的正确性判断下面逐层展开。2. FlashAttention的底层机制分块、SRAM、在线softmax2.1 标准attention的瓶颈在于HBM读写要理解FlashAttention先看标准attention在GPU上的数据流动。Q、K、V都存在显存里显存在英伟达的术语里叫HBM。标准attention大致分四步读Q和K计算S QK^TS是n×n的。把这个S写回HBM。softmax需要整行的最大值和指数和所以要把S重新读回来算完P再写回HBM。最后读P和V计算PV。这个过程中n×n的S矩阵被反复读写。GPU上常见的瓶颈往往不是浮点算力而是数据搬运速度。HBM带宽和片上SRAM带宽差一个数量级以上。FlashAttention论文里最核心的一张图就是展示attention耗时大头在HBM读写而不是FLOPs。2.2 分块计算如何减少数据搬运FlashAttention的思路是让注意力计算尽量保持在片上SRAM里不要频繁进出显存。做法是把Q、K、V切成小块。比如每块大小是128×128。一次循环只处理一个Query块然后遍历对应的Key块和Value块在SRAM里算分数、做softmax、累积输出。整个过程里n×n的分数矩阵从来没有完整生成过只有块级的临时矩阵。这带来两个直接好处HBM访问次数大幅减少。每个K/V块只被加载一次而不是每次计算都要从HBM重新拉。内存占用从O(n²)降到O(n)。不需要为整个注意力矩阵分配显存