FlashAttention与滑动窗口注意力融合:加速长文本prefill

发布时间:2026/9/1 18:17:10
FlashAttention与滑动窗口注意力融合:加速长文本prefill 如果正在做长文本大模型的推理优化大概率会遇到这样一个场景prompt 已经到几千 tokenprefill 阶段 GPU 算力拉满第一个 token 却迟迟出不来。为了把上下文做长很多人会引入滑动窗口注意力为了把计算提速又会想到 FlashAttention。但这两个方案各自都好理解一旦要叠加使用很多人就卡住了直接给 FlashAttention 加一个窗口 mask真的能加速吗先给一个判断FlashAttention 优化的是计算调度它不改变注意力的数学定义滑动窗口注意力优化的是计算结构它把稠密的注意力依赖变成带状稀疏依赖。这两者不是一个层面的东西所以“给 FlashAttention 加一个 mask”离真正的加速还差得很远。只有把窗口稀疏性翻译成 FlashAttention 的 block 级跳过才能同时拿到两者的收益这也是本文标题里那个问题的核心答案。这篇文章会从 prefill 的瓶颈讲起先把 FlashAttention 的原理、滑动窗口注意力的特性拆开再给出一份可以照着实现的分块窗口注意力代码最后讲清楚验证方法、性能分析思路和工程落地时容易踩的坑。整个过程不依赖特定推理框架理解之后也可以迁移到自己的项目中。1. 为什么这个问题值得单独拿出来写先明确一下背景LLM 生成文本时推理过程被分成 prefill 和 decode 两个阶段。prefill 阶段做的事情是把用户输入的一整段 prompt 一次性喂给模型并行计算所有位置的注意力同时生成 KV Cache输出第一个 token。这个阶段是计算密集型的因为要处理完整的 prompt而且整个过程可以并行。问题在于如果使用标准的 full attention位置数量 n 对应的注意力矩阵是 n×n计算量和显存占用都是 O(n²)。prompt 越长prefill 就越慢直观表现就是“第一个 token 等了特别久”。滑动窗口注意力的思路很直接每个 token 不再和全部历史 token 做注意力而只和最近 W 个 token 做注意力。这样每个 token 只需要关注固定大小的窗口整体复杂度从 O(n²) 降到 O(nW)。当 W 远小于 n 时节省是数量级的。Longformer、BigBird 以及一些现代大模型都采用了窗口注意力或窗口加全局 token 的混合设计。FlashAttention 是另一条优化路径。它面对的还是 full attention但通过 IO 感知的分块计算把注意力计算过程中的大量 HBM 往返读写降到最低跑得更快。很多人因此误以为 FlashAttention 可以随便和稀疏注意力叠加实际上没有那么简单。真实情况是FlashAttention 的 tiling 调度按固定 block 切分滑动窗口的稀疏模式如果不在 block 粒度对齐你会陷入两种糟糕情况——要么算了不该算的 block要么为了跳过 block 付出了更高的索引和 mask 成本。所以理解“FlashAttention 如何加速滑动窗口注意力 prefill”本质上是理解两个经典优化如何正确组合。这个问题适合以下读者做 LLM 推理服务优化、做长文本工具链、写训练或推理 kernel或者只是想把注意力机制底层逻辑搞清楚的人。2. FlashAttention 到底加速了什么从内存说起FlashAttention 的加速效果经常被一句“减少显存占用”带过但真正关键的是它改变了注意力计算对 HBM 的访问模式。先看朴素注意力在 GPU 上的执行路径。给定 Q、K、V流程是计算 S QK^T得到 n×n 的分数矩阵。对 S 做 scale 和 softmax得到概率矩阵 P。计算 O PV。每一步都可能把中间矩阵写回显存也就是 HBM。GPU 里有很大的 HBM但读写速度远低于芯片上的 SRAM。注意力计算的核心问题不是“算力不够”而是“数据搬运太慢”。n×n 的中间矩阵越大HBM 的读写压力越大。这也是为什么很多 long context 任务在算力看起来充足的情况下prefill 时间仍然随序列长度急剧上升。FlashAttention 的改进思路是让数据尽可能在 SRAM 里完成计算。它把 Q、K、V 都切成小块每次只在 SRAM 中加载一个小的 block完成局部 QK^T、softmax、PV 计算再累加结果全程不写回 n×n 的中间矩阵。这带来一个看似矛盾的问题softmax 需要对整行计算最大值和归一化项分块之后每块只有局部信息怎么办答案是 online softmax。它维护两个运行状态当前行的最大值 m以及归一化项的指数和 l。每次处理一个新的列块时先用新的块计算局部最大值再更新全局 m用 rescale 因子把之前累加的结果调整到新的数值范围最后累加新的概率和值。这样一来虽然结果和标准 softmax 完全一样但中间矩阵 P 从头到尾都不需要完整存在。可以这样对比对比维度朴素 AttentionFlashAttention中间矩阵 S/P需要完整写回 HBM在 SRAM 中局部计算不写回计算调度矩阵乘、softmax、矩阵乘分开执行融合到分块循环中softmax 处理整行先算 max 和 sumonline softmax 动态更新主要瓶颈HBM 带宽算力访存开销大幅下降理解到这一层会发现 FlashAttention 本身不改变注意力语义也不把计算量减小。它做的是让同样的计算变得更“顺”减少数据搬移的代价。所以它适合用来加速 full attention也适合用来加速窗口注意力但前提是窗口的稀疏方式不能破坏 tiling 的访存收益。第 2 章的核心结论是FlashAttention 是 IO 调度层面的优化而不是算法结构层面的优化。我们要把滑动窗口注意力嵌入 FlashAttention需要做的是在它的 tiling 调度上表达窗口的 block 级稀疏。3. 滑动窗口注意力从算法稀疏到内存稀疏滑动窗口注意力的公式不复杂。对于位置 i注意力只允许关注满足 j ≥ i - W 的 key/value这里 W 是窗口大小。在实际实现中通常会同时加上 causal mask所以 j 还要满足 j ≤ i。于是位置 i 的注意力范围是 [i-W, i]。复杂度上每个位置只和 W 个位置做注意力整体计算量是 O(nW)。如果 W 固定那就和序列长度 n 保持线性关系。这也是为什么窗口注意力经常出现在长文档、多轮对话这类场景中——上下文可以无限增长但单次注意力的计算量不会随之无限膨胀。不过“算法稀疏”不等于“内存稀疏”。如果你直接在 PyTorch 里用一个attn_mask把窗口外的位置填成-inf底层仍然会执行完整的 QK^T 运算稀疏 mask 只是让 softmax 之后的值趋向 0。也就是说计算量没有真正降下来只是结果看起来像窗口注意力。这种做法在序列很短时没问题在序列很长时prefill 依然会撞上 O(n²) 的计算墙。更麻烦的是稀疏 mask 对访存的影响。GPU 的 kernel 喜欢连续、规整的数据访问。窗口注意力的有效区域是一条斜带状如果按元素粒度控制每个 warp 处理的相邻位置可能落在完全不同的列范围导致访存不连续缓存命中率下降。结果可能是虽然理论上少算了很多 FLOPs但实际运行时间并没有按预期下降。那么什么时候“算法稀疏”会真正变成“内存稀疏”答案是把稀疏模式对齐到固定 block 粒度。FlashAttention 的分块循环天然按 block 遍历如果窗口能恰好覆盖整数个 block就可以在 block 级跳过不需要的列块访存路径变得规整计算也真正跳过。这正是下一章要展开的核心实现。第 3 章的核心结论是滑动窗口注意力省掉的是计算量但只有把窗口翻译成 block 级稀疏才能同时省掉 HBM 访问真正缩短 prefill 时间。4. 环境准备与前置知识本文后续示例使用 Python 和 PyTorch 风格实现目的是演示算法逻辑而不是提供一个生产级 kernel。真实落地时我们还需要把它改写成融合 kernel例如用 Triton 或者 CUDA 实现。版本信息以实际项目为准这里不限定某一个具体版本。建议的验证环境Python 3.8 以上版本。PyTorch 2.x这样可以使用torch.nn.functional.scaled_dot_product_attention做对照实验前提是硬件和 CUDA 版本满足条件。一块 NVIDIA GPU用于后续性能验证如果只是验证数值正确性CPU 也能跑通。CUDA 工具链和 Triton方便后续把窗口逻辑扩展到融合 kernel。先做一个最小环境检查确认 PyTorch 和 GPU 可用import torch print(PyTorch version:, torch.__version__) print(CUDA available:, torch.cuda.is_available()) print(GPU name:, torch.cuda.get_device_name(0) if torch.cuda.is_available() else CPU)5. 核心实现用 FlashAttention 的思路加速窗口 prefill5.1 窗口对齐到块最关键的一步FlashAttention 的分块逻辑中一个 block 是大小为 B×B 的矩阵块。滑动窗口的有效范围是斜带状如果要让一个 block 要么完全在窗口内、要么完全在窗口外最稳妥的做法是把窗口大小 W 设置成 block size B 的整数倍。假设 B 32W 64。那么对于任意行块 i需要关注的列块范围是 [i - 2, i]在不考虑边界的情况下。也就是说每个行块只需要和最多 3 个列块做注意力计算而不是和所有行块。当序列长度 n 远大于 W 时窗口外的列块可以被直接跳过。如果 W 不是 B 的整数倍会出现一个 block 内部部分有效、部分无效的情况这时候就必须做元素级 mask。元素级 mask 在 softmax 之前填入-inf会增加一次数值操作而且破坏了 block 整体跳过的能力。所以工程上建议统一约定窗口大小一律取 block size 的整数倍。5.2 分块遍历与 block 级跳过现在把 FlashAttention 的行块循环改造成窗口版本。伪代码如下for row_b in range(n_blocks): row_start row_b * block_size row_end min(row_start block_size, seq_len) # 窗口左侧边界对应的列块 index left_block max(0, row_b - window_blocks) for col_b in range(left_block, row_b 1): col_start col_b * block_size col_end min(col_start block_size, seq_len) # 只在 block 层面做 QK^T、online softmax、PV 累加这里window_blocks W / B。外层循环仍然按行块遍历内层循环只遍历当前行块窗口范围内的列块。这样窗口外的列块完全不进入 QK^T 计算省 FLOPs。访存范围集中在窗口附近的 KV 块HBM 访问路径更规整。每个列块内部又可以复用 FlashAttention 的 tiling 思路。如果不想依赖“窗口是块整数倍”这个前提保留元素级 mask 做边界处理也可以但性能会打折扣。下面的完整示例同时演示了 block 级跳过和边界 mask 的写法保证任意 W 下结果都正确。5.3 online softmax 处理窗口块即使只遍历窗口内的列块也不能直接先算完整行的 softmax因为完整行的分数矩阵并没有被一次性算出来。我们仍然需要在遍历列块时维护m当前行处理到的最大的注意力分数。l当前行归一化项的指数和。acc当前累计的 PV 结果。每处理一个列块先算局部最大值更新全局 m然后 rescale 之前的累计值m_new torch.maximum(m, s.max(dim-1).values) p torch.exp(s - m_new.unsqueeze(-1)) l l * torch.exp(m - m_new) p.sum(dim-1) acc acc * torch.exp((m - m_new).unsqueeze(-1)) p v_block m m_new这和标准 FlashAttention 的 online softmax 完全一致。区别在于标准版本内层遍历所有列块窗口版本只遍历窗口覆盖到的列块。5.4 完整示例代码与运行验证下面是一个可直接运行的教学版本。第一个函数是朴素窗口注意力用来验证正确性第二个函数是窗口版 FlashAttention 分块实现最后一段脚本会跑随机数据对比两者输出。import torch import torch.nn.functional as F def naive_window_attention(q, k, v, window_size): 滑动窗口注意力的朴素实现用于验证正确性。 q/k/v: [seq_len, head_dim] 窗口语义位置 i 只关注 [i - window_size, i] 范围内的 key。 seq_len q.shape[0] out torch.zeros_like(q) scale q.shape[-1] ** 0.5 for i in range(seq_len): start max(0, i - window_size) scores q[i] k[start:i 1].T / scale weights torch.softmax(scores, dim-1) out[i] weights v[start:i 1] return out def flash_window_attention(q, k, v, block_size32, window_size64): 教学用分块实现窗口 FlashAttention 核心逻辑。 说明这里用 PyTorch 循环演示 tiling online softmax 真实高性能版本需要写成 CUDA/Triton 融合 kernel。 seq_len, dim q.shape scale dim ** 0.5 out torch.zeros_like(q) n_blocks (seq_len block_size - 1) // block_size window_blocks (window_size block_size - 1) // block_size for row_b in range(n_blocks): row_start row_b * block_size row_end min(row_start block_size, seq_len) q_block q[row_start:row_end] acc torch.zeros(row_end - row_start, dim) m torch.full((row_end - row_start,), float(-inf)) l torch.zeros(row_end - row_start) left_block max(0, row_b - window_blocks) for col_b in range(left_block, row_b 1): col_start col_b * block_size col_end min(col_start block_size, seq_len) s q_block k[col_start:col_end].T / scale # 边界 mask保证 causal 且窗口外位置为 -inf mask_row torch.arange(row_start, row_end).unsqueeze(1) mask_col torch.arange(col_start, col_end).unsqueeze(0) valid (mask_col mask_row) ((mask_row - mask_col) window_size) s s.masked_fill(~valid, float(-inf)) # online softmax m_new torch.maximum(m, s.max(dim-1).values) p torch.exp(s - m_new.unsqueeze(-1)) l l * torch.exp(m - m_new) p.sum(dim-1) acc acc * torch.exp((m - m_new).unsqueeze(-1)) p v[col_start:col_end] m m_new out[row_start:row_end] acc / l.unsqueeze(-1) return out运行验证脚本torch.manual_seed(0) seq_len 128 dim 16 q torch.randn(seq_len, dim) k torch.randn(seq_len, dim) v torch.randn(seq_len, dim) out_ref naive_window_attention(q, k, v, window_size64) out_flash flash_window_attention(q, k, v, block_size32, window_size64) print(max abs diff:, (out_ref - out_flash).abs().max().item())预期输出会是一个很小的浮点差异通常在1e-6量级。之所以在 fp32 下也有一些微小的差异是因为 online softmax 的 rescale 每一块都有一次指数运算浮点数累积顺序不同。如果使用 fp16 或 bf16差异会稍微变大这是正常现象。验证通过说明分块实现从头到尾和朴素窗口注意力的数学语义是一致的。这个 forward 逻辑放在 prefill 场景中就是一次性编码全程 prompt 的过程。这里要强调一点上面这份代码是正确性演示不是性能演示。用 PyTorch 的 Python 循环模拟 tiling只是为了把 FlashAttention 的算法意图说清楚。真正的加速来自把这段线性循环写进融合 kernel让 block 级跳过真正发生在 GPU 调度层面。6. 从 prefill 到 decode窗口注意力对 KV Cache 的连锁影响prefill 完成后注意力计算过程中得到的 K、V 会被保存为 KV Cache供 decode 阶段使用。窗口注意力对 KV Cache 的影响比“加速 prefill”本身更深远。在 full attention 中每个新 token 生成时都要读取全部历史 KVdecode 阶段每步的访存成本随上下文长度线性上升。窗口注意力提出了一种更经济的方案位置 i 的 K、V 只需要服务未来 W 个位置。超过 W 之后它不会再被任何后续 token 读取。这意味着两件事。第一KV Cache 可以做成滑动淘汰的。最常用的实现是环形缓冲存储一个固定长度为 W 的 KV 区域新 token 的 KV 写入头部把最老的 KV 覆盖掉。内存占用从 O(n) 降到 O(W)上下文长度不再直接决定 KV Cache 增长。这个设计对长文档服务非常关键因为 KV Cache 往往比模型权重本身占用的显存还大。第二decode 阶段的注意力计算范围也被固定。每次生成新 token只需要读取最近 W 个 token 的 KV而不是全部历史 KV。这一步直接降低了每 token 延迟的后半段时间。所以滑动窗口注意力从来不只是“让 prefill 更快”的工具。它从 prefill 阶段开始减少计算到 decode 阶段减少 KV 读取同时控制显存占用。理解这条链路才会明白为什么很多长文本部署方案会把窗口注意力和 KV Cache 淘汰一起做。有一个需要警惕的点如果模型结构里除了窗口注意力之外还包含全局 token比如某些位置始终能看到全部上下文KV Cache 就不能简单做全量淘汰。否则该读的全局信息丢了输出质量会下降。这类混合结构需要分别处理窗口 KV 和全局 KV窗口部分滚动淘汰全局部分长期保留。7. 性能分析思路别只盯着 FLOPs很多人看到窗口注意力的复杂度是 O(nW)第一反应就是“FLOPs 降了这么多肯定更快”。这个判断在数学上成立在实际工程中却经常失灵。原因在于GPU 上的实际耗时取决于三件事计算量、访存量、kernel 的调度效率。只降计算量但访存路径混乱可能得不偿失。比如在 PyTorch 里直接构造一个 n×n 的 bool mask 传给 attention虽然理论上只关注窗口内但 mask 矩阵本身是 n×n额外显存和访存开销已经很大性能甚至会低于 full attention。有效的性能分析应该区分三种实现形态实现形态计算量中间矩阵/访存是否获得 FlashAttention 收益朴素 dense maskO(n²) 实际仍全算n×n mask 和中间矩阵否稀疏 mask FlashAttention 库O(nW) 但未做 block 对齐依赖 mask 实现可能仍有额外开销部分窗口 block 级跳过 融合 kernelO(nW)只访问窗口内 KV block是当序列长度较短时第三种形态未必比 full attention 成熟 FlashAttention 快因为 block 索引判断、循环控制和 kernel launch 也有开销。通常当 n 远大于 W 时加速效果才明显。真正的验证方法是做一组系统对照实验选择一批相同输入固定 batch size 和序列长度。分别跑 full attention 和窗口 block 版。记录 prefill 耗时、KV Cache 峰值显存、首 token 延迟、decode 阶段每 token 延迟。对比输出误差确认数值一致性。如果希望更准确地定位瓶颈用 profiling 工具观察各 kernel 的耗时占比和 HBM 读写量会比只看总时间更有参考价值。8. 常见问题与排查方法问题现象可能原因排查方式解决方案结果出现 NaN 或 Infmask 中使用了-inf在 fp16 下可能导致指数溢出检查 mask 构造和 softmax 前的数值范围使用较大负数如-65500或采用 block 级跳过避免元素级 mask窗口外 block 跳过后结果和朴素实现不一致边界条件判断错误可能把窗口边缘的 block 错误跳过先跑第 5 节中的数值对比脚本用window_blocks ceil(W / B)处理不整除情况保留边界 mask序列长度不是 block_size 的整数倍最后一块不足 B 时行块和列块长度不等检查最后一个 block 的row_end和col_end使用 min 截断并确保 mask 按实际长度生成使用成熟库的 attention_mask 发现没有真正加速库内部仍按 dense 方式处理 mask未做 block 级跳过查看 kernel 耗时比对 HBM 读写量改用 block-sparse 内核或自己实现分块循环KV Cache 淘汰后输出质量明显下降模型结构包含全局 attention 或其他对窗口外 token 的依赖检查模型 config 中的 attention 类型对全局 token 单独保留 KV不做滚动淘汰CPU 上运行示例很慢示例代码是 Python 循环模拟仅用于正确性验证不要用该版本测性能将逻辑改写成 Triton/CUDA kernel再在 GPU 上测耗时9. 工程最佳实践与落地建议先把最实用的几个结论写在前头。第一窗口大小设为 block size 的整数倍。这能保证 block 级边界对齐避免元素级 mask是“算法稀疏”变成“内存稀疏”的前提。如果你用的 block size 是 32窗口就取 512、1024 这类值。第二block 索引提前一次性算好。不要在 kernel 内部重复判断每个 block 是否属于窗口而是先把每个行块的left_block和right_block预计算出来作为索引范围传入。省掉的不仅是计算还有分支判断带来的 warp divergence。第三先验证数值正确性再优化性能。用朴素窗口注意力做基准对比分块实现的最大绝对误差。这一步确认语义正确后再接入真实长文本数据才不至于把算力浪费在一个“看起来很对但结果不对”的实现上。第四优先复用成熟组件。PyTorch 较新版本的scaled_dot_product_attention在满足条件时底层会走到融合 kernel。如果它不能满足窗口 block 跳过的需求可以再看推理框架里是否已经有 varlen 或 block-sparse attention 支持。都满足不了时再考虑用 Triton 写一个专门的窗口 FlashAttention kernel。第五做好监控指标。在 prefill 阶段重点看首 token 延迟和峰值显存在 decode 阶段重点看每 token 延迟和 KV Cache 大小。任何局部优化都要放到整条推理链路里评估防止“prefill 快了decode 反而慢了”这种顾此失彼的情况。第六长序列测试要覆盖边界条件包括窗口小于 block size、序列长度略大于窗口、序列长度是 block size 的整数倍等。很多看似微小的问题只会在边界条件下暴露。10. 总结与后续学习方向回到标题里的问题FlashAttention 加速滑动窗口注意力 prefill核心不是给注意力加一个稀疏开关而是把窗口稀疏性翻译成 block 级跳过再和 FlashAttention 的 tiling 调度融合。这条路径的关键点有三个FlashAttention 本身不改变注意力数学定义只改变计算调度和访存方式滑动窗口注意力把注意力计算从 O(n²) 变成 O(nW)但要真正省时间必须做 block 级跳过