
当输入序列从 2K 涨到 32K很多模型在 prefill 阶段明显变慢。做过长文本推理的人大概率会想到用滑动窗口注意力每个 token 只关注最近的 W 个 token把 attention 的计算量从 O(N^2) 降到 O(N·W)。但真把它和 FlashAttention 放到一起时问题并不只是“加一个 mask”那么简单。滑动窗口注意力的稀疏模式、prefill 阶段的 I/O 特点以及 kernel 的分块方式决定了你是在真正加速还是在用更复杂的方式绕远路。这其实是一个很值得拆开的问题FlashAttention 为什么能加速标准注意力这个机制到了滑动窗口注意力上哪些优势还在哪些优势会被削弱如果你正准备把长上下文模型部署到生产环境或者正在做训练/推理性能优化这篇文章可以当作一个从原理到落地的检查清单。1. prefill 真正复杂的地方不是计算量而是中间矩阵1.1 prefill 和 decode 是两条优化路径模型推理按阶段分可以分成 prefill 和 decode。prefill 发生在你输入完整 prompt 时模型一次性处理所有 token并为每个 token 生成 KV cache。decode 是之后一个一个生成新 token 的阶段。两者对延迟的敏感度完全不同prefill 是“一次处理一大批 token”decode 是“一次只处理一个 token”。所以优化 prefill重点不是让单个 token 变快而是让整个 prompt 的处理过程在并行度和内存访问上更高效。在标准 attention 实现里prefill 会为每个 query 位置计算对所有 key 位置的注意力分数形成一个 N×N 的注意力矩阵。这个大矩阵往往不会全部驻留在寄存器里而是要写回 HBM高带宽内存下一次读取时再读回来。这个“写入再读回”的过程才是 attention 在长序列下变慢的主要原因。简单说瓶颈不只在 FLOPs还在内存带宽。FlashAttention 的出发点正是这个不把 N×N 的注意力矩阵完整落地到显存而是分成小块计算在块内用 online softmax 更新统计量最后只把输出写回。这样避免了中间矩阵的重复写入和读取。对于这一步标准 attention 和滑动窗口 attention 面临的问题是相同的。1.2 滑动窗口注意力让矩阵变“瘦”却不等于更便宜滑动窗口注意力也叫 sliding window attention通常约定每个位置 i 只 attend 到 i-W 到 i 之间的 key。如果配合 causal mask就大约只保留主对角线附近的一条“带”真正的注意力矩阵从全矩阵变成带状稀疏矩阵。理论上合法注意力对的数量从 N^2 减少到大约 N·W。但你如果只是用 PyTorch 写一个 mask比如attn_score.masked_fill(mask, -inf)矩阵乘法仍然会把所有位置都算一遍mask 只是把最终结果遮掉。计算量没有变少反而多了一次 mask 操作和额外的内存访问。更糟的是滑动窗口会让 attention 矩阵看起来有规律但普通矩阵乘法不会利用这种带状规律。很多第一次做局部注意力的人发现加了窗口后速度没变甚至更慢就是这个原因你只是让“不该看的地方变暗”并没有让“不该算的地方不计算”。这就引出了 FlashAttention 的另一个价值它天然按块组织计算可以在块级别跳过那些落在窗口之外的 key 块。注意前提是 kernel 真的做了这个跳过逻辑而不是简单地把 mask 传进去。2. FlashAttention 加速滑动窗口注意力的核心逻辑2.1 把完整矩阵拆成小块是起点FlashAttention 的思路可以简化成三步把 Q、K、V 都切成 block对每个 query block遍历它需要的 key/value block在块内计算 attention 分数并累加输出。因为每个 block 足够小可以放进 shared memory 或寄存器不需要把完整注意力矩阵写回显存。这是它提升 IO 效率的根因。对标准 attention 来说每个 query block 通常要遍历所有 key block即使有 causal mask也只遍历到当前 query block 对应的位置附近。对于滑动窗口这个遍历范围可以进一步缩小每个 query block 只需要遍历与它窗口有交集的 key block。如果窗口 W 远小于序列长度 N那么每个 query block 需要访问的 key block 数量就大致等于 W 的覆盖宽度而不是 N。这样一来整体 FLOPs、HBM 访问量、以及中间 softmax 统计量都会随窗口大小线性下降。这是“块级稀疏”带来的收益。2.2 在线 softmax 让分块结果能安全合并但分块处理 attention 有一个数学上的麻烦softmax 有全局归一化项直接按块算分数不能简单加总。FlashAttention 的做法是维护每个 block 的局部最大值 m 和局部指数和 l处理完一个 key block 后用新的最大值对已经累加的输出做 rescale。这个机制对滑动窗口同样适用。你只需要在 block 内部把不在窗口范围内的分数设置为负无穷。因为负无穷参与 softmax 后指数为 0不会影响局部最大值和指数和。不过要注意-inf 本身在数值上需要小心处理很多实现会用很小的负值或者通过 block 级跳过避免进入计算。这里有一个关键区分块级跳过处理的是“整个 block 都不在窗口内”的情况块内 mask 处理的是“block 部分在窗口内”的情况。如果窗口边界正好落在 block 中间你不能跳过整个 block只能通过 mask 把不该出现的位置遮掉。因此block 大小和窗口大小之间的关系会直接影响实际收益。2.3 跳过与窗口无关的块才是滑动窗口真正提速的地方很多人把 FlashAttention 加在滑动窗口注意力上只是把 mask 换成窗口 mask结果发现收益很小。问题出在 kernel 没有减少 block 迭代次数。一个有效的滑动窗口 FlashAttention kernel在外层循环上就要计算出当前 query block 的 key 范围然后只遍历落在范围内的 key block。这个“跳过”发生在 kernel 的 block 遍历逻辑里不是藏在 mask 里。用本文前段的话说你要让“不该算的地方根本不算”。具体范围是这样的假设窗口只关心当前 query 前面的 W 个 key即因果滑动那么对于一个从 qs 到 qe 的 query block它需要访问的 key 位置大致从 max(0, qs - W) 到 qe。如果 W 比序列长度小一个数量级这个范围远比把整个序列遍历一遍要短。代码层面这个范围就是一个 for 循环的起点和终点。一旦确定好窗口注意力的主要计算量就只集中在带状区域内。这才是 FlashAttention 加速滑动窗口注意力 prefill 的完整逻辑。3. 从标准到滑动窗口一个可落地的 kernel 计算框架3.1 外层循环query block 能看到的 key block 范围下面给出一个针对因果滑动窗口的 FlashAttention 外层遍历结构。注意这不是某个库的完整源码更像是一个理解原理的“伪代码映射”。# 示意代码滑动窗口 FlashAttention 的 block 遍历逻辑 # seq_len: 输入长度 # W: 滑动窗口大小例如 1024 # Q_BLOCK, KV_BLOCK: 分块大小例如 128 for qs in range(0, seq_len, Q_BLOCK): qe min(qs Q_BLOCK, seq_len) # 当前 query block 可能关注的 key 范围 # 因果滑动key 需要在 [max(0, qs - W), qe) 内 k_start max(0, qs - W) k_end qe # 因果下key 不需要超过当前 query block 的末尾 # 遍历这个范围内的 kv block for ks in range(k_start, k_end, KV_BLOCK): ke min(ks KV_BLOCK, k_end) # 1. 从 HBM 加载 Q_block, K_block, V_block # 2. 计算 scores Q_block K_block^T # 3. 对窗口外的位置设置 -inf # 4. 分块内做 softmax 统计更新 output这里我特意把 k_end 设置成 qe而不是 seq_len。原因是 causal mask 已经保证 query 不会 attend 到未来的 key。如果只做双向滑动窗口不叠加因果那么 k_end 可以到 min(seq_len, qe W)。实际应用里语言模型通常叠加 causal mask所以“因果窗口”是更常见的组合。另外注意 k_start 计算对于 query block 内最早的那个 query索引 qs它的合法 key 从 qs - W 开始对 block 内最晚的 query从 qe - 1 - W 开始。为了保证整个 block 都不会漏掉合法 keyk_start 应该取整个 block 的最小值也就是 qs - W。这是很多初学者容易错的地方如果把 k_start 设成 qe - 1 - W前面一部分 query 的窗口会被截断。3.2 mask 应该怎么设计才不会丢掉边界即使外层循环只遍历有交集的 key blockblock 内部仍然可能出现两类“非法”位置当前 query 未来的 keycausal mask。距离超过 W 的 keywindow mask。你可以把两者的交集合成一个 mask也可以分开处理。在 block 遍历范围内causal mask 和 window mask 叠加后有效区域是一个左下角的梯形/带状区。如果使用 -inf 填充一定要保持 mask 的 dtype 与 score 一致并且在没有使用 fp16 混合精度时避免 -inf 被转成 NaN。一个更稳妥的做法是在 block 内部生成一个布尔 mask对于合法位置保持 score非法位置设为极小负值。对 fp16/bf16 来说-inf 有时会出现异常常见实现会选择-65500或-3.4e38这类足够小的值。具体数值可以结合实现经验来调但重点是这个值要远小于正常 score 的绝对值不能影响 softmax 结果。如果窗口比较大比如 W 是 Q_BLOCKKV_BLOCK 的整数倍block 边界和窗口边界错位的概率会降低很多 block 可以直接判定为完全在窗口内连 mask 都不需要生成。这也是为什么在 FlashAttention 场景里block size 往往要参与调优而不是用固定的 64。3.3 结合 causal、window、padding 时的组合顺序当输入是多个样本 padding 成一个 batch 时padding 位置也会占用窗口范围。如果某个 key 是 padding token它不该被 attend 到同时padding 的 query 也不该产生有意义的输出。这个必须在 mask 里一起处理。实际组合顺序可以这样先确定窗口范围正常 position 在 [qs - W, qe) 内的 key 可以被读取。再叠加 causal maskkey index query index 才合法。最后叠加 padding mask非 padding 位置才合法。如果某个 key block 完全超出这三者的合法范围直接跳过如果部分重叠块内用 mask 处理。这个顺序不是唯一答案但它比较直观。先在 block 粒度筛掉大量无关 key再在 block 内部处理边界比一开始就对每个 score 做 mask 要高效得多。实际 kernel 里还要考虑如果窗口范围很小kv block 数量很少就应该把 Q_BLOCK 调大减少 query block 的启动和遍历开销如果窗口很大接近 N则更应该把 Q_BLOCK 和 KV_BLOCK 调小减少块内无效计算。4. 性能收益能有多大边界又在哪里4.1 理论上节省的 FLOPs 和内存访问假设序列长度 N4096窗口 W512。标准 causal attention 需要计算约 N^2/2 8.4M 个 score 对因果滑动窗口后大约为 N*W 2.1M 个少了一个量级。如果换成 N32768、W1024标准计算量约 536M窗口计算量约 33M差距更明显。这里说的是“合理的 score 数量”不一定是实际 kernel 里计算的累加次数因为 block 边界的部分无效元素会产生额外计算但趋势是清晰的。对内存访问来说FlashAttention 的收益也体现在 KV 的读取量上。标准 FlashAttention 处理一个 query block理论上也要访问所有 key/value block滑动窗口版本只访问窗口内的 key/value block。所以序列越长、窗口相对越小HBM 访问量的节省越明显。这也是很多长文本模型在 prefill 阶段使用滑动窗口注意力的原因不是 FLOPs 降了而是整个计算链路上的数据搬运量降了。4.2 为什么短序列、小窗口可能反而更慢FlashAttention 的加速依赖两个前提block 粒度足够大以及被跳过的 block 数量足够多。如果序列只有 512窗口却有 128那么每个 query block 仍然要访问不少 kv block跳跃收益不明显。再加上 kernel 启动、mask 生成、block 遍历逻辑这些固定开销最终可能比一个朴素 mask 实现还慢。还有一类情况更尴尬窗口大小和 block size 不匹配比如 W100Q_BLOCK128KV_BLOCK128。这时每个 query block 几乎总要碰到两个 kv block 边界第二个 kv block 里只有很少的合法位置。计算一个 128×128 的 score block最后可能只保留十几列大量算力被浪费在无效 score 上。所以不要只看“窗口是 100”就觉得工作量是 N×100实际 kernel 的迭代次数可能是 N×2 个 kv block而不是 N×1 个。4.3 区分“计算被跳过”和“内存访问被跳过”有一个常见误解只要 FlashAttention kernel 里跳过了某段 for 循环那么相关的 K/V 一定不会被读取。实际上如果循环是从某个 k_start 到 k_end那么被跳过的 key block 确实不读。但如果 block 内部部分被 maskK/V 是被完整加载的只是部分 score 被 mask。所以性能收益取决于你“跳过”了多少个 block而不是“mask”了多少个元素。实践中可以这样判断把窗口 W 缩小到原来的 1/2如果耗时没有明显下降说明 block 内部 mask 或固定开销占了大头你需要调整 block size或者考虑使用更细粒度的稀疏 kernel。如果耗时确实明显下降说明你的遍历逻辑已经尽量把无关 block 挡在门外了。这个经验可以用于快速定位滑动窗口 FlashAttention 是否真的在“加速”。5. 落地中的建议步骤和典型坑点5.1 先跑一个小规模正确性测试不要急着上长序列一个常见习惯是拿到 kernel 后直接拿去跑 32K 长文本看到显存下降就觉得优化成功。但注意力类 kernel 最容易出的问题不是显存而是 mask 边界错误导致的结果错位。最好先构造一个小用例比如 N32W8Q_BLOCK8KV_BLOCK8用参考实现标准 PyTorch causalsliding window attention对比数值。因为 N 和 W 都很小你可以逐个检查位置 i 到底 attend 了哪些 key确认窗口边界、causal mask、padding mask 组合正确后再放大到真实序列。小规模验证时建议同时检查attention 输出与参考实现的绝对误差/相对误差每个 token 的实际注意力权重是否落在窗口内第一个 block 和最后一个 block 的边界行为使用 fp16/bf16 时是否有 NaN 或 inf。这些检查点其实都不用写很多代码一个 pytest 测试或脚本即可完成。但很多人跳过这一步直接进性能调优后面问题找起来会非常痛苦。5.2 性能基准要分别测 prefill、端到端和显存如果要给团队或项目做选型不建议只测一个总耗时。prefill 的时间、decode 的每 token 延迟、峰值显存、吞吐量这四类指标反映不同的问题。滑动窗口FlashAttention 主要优化 prefill 阶段和显存占用decode 阶段是否受益主要看实现是否在读取 KV 时也利用了窗口。如果 decode kernel 仍然把全部 KV 读进来prefill 再快端到端 token 生成速度也未必提升。基准测试还要固定以下条件序列长度、窗口大小、batch size、block size、mask 实现方式、是否使用因果、硬件型号。否则很难判断一个改动是变好还是变坏。我会建议每次只改一个变量。比如先固定 batch1调 block size再固定 block size调窗口大小。5.3 常见问题排查链路如果在调试“FlashAttention滑动窗口”时出现问题可以按这条链路查看现象是 loss 不收敛、输出 NaN、速度没变还是显存爆掉不同现象指向不同原因。看 mask先确认窗口范围有没有“往前多算一格”或“少算一格”。把窗口边界打印出来检查。看遍历范围确认外层 for 循环的 k_start/k_end 是否覆盖了所有合法 key block有没有把 padding block 也包进来。看 block 大小如果 W 与 block size 不是倍数关系边界 block 会浪费大量算力考虑调整 Q_BLOCK/KV_BLOCK。看数值稳定性在线 softmax 对 -inf 的依赖较重极小负值要足够小同时注意 fp16 下指数运算的上下溢出。看硬件信息某些 GPU 架构对 shared memory 大小有限制block 过大会导致 kernel 启动失败或回退到 fallback 实现。这条链路不保证一次找出问题但能帮你从“感觉不对”走向“定位到具体层”。很多 attention kernel 的问题最后都出在 mask 或遍历范围而不是数学公式。6. 判断框架什么时候值得用“FlashAttention滑动窗口”6.1 四步判断法如果要把这个方案放进自己的项目我建议先走四步第一步确认你的模型是否真的需要局部注意力。这取决于任务类型长文本检索、文档摘要、代码生成可能更需要全局依赖滑动窗口不一定合适。第二步先确认标准 attention或标准 FlashAttention到底慢在哪。如果序列长度还没超过 2K瓶颈可能是模型算子本身而不是注意力矩阵规模。第三步小规模验证窗口注意力对模型效果的影响。窗口太小会导致长距离信息丢失如果模型本身没有局部归纳偏置效果可能明显下降。第四步再考虑 kernel 层面怎么加速。这时你已经知道窗口可以接受剩下的才是 FlashAttention 的遍历范围、block size、mask 和数值稳定性问题。这个方法帮你把“要不要换模型结构”和“怎么把结构变快”分开。很多人把顺序搞反先优化 kernel才想起来验证效果最后发现窗口设置毁了模型。6.2 适用场景与不适用场景适用场景长上下文的训练或 prefill序列长度远大于窗口大小。内存受限完整 attention 矩阵或完整 KV cache 放不下。任务本身对近邻依赖强窗口内信息基本够用。已有标准 FlashAttention 或相关 kernel 基础想扩展成窗口形式。不适用场景短序列或小 batchkernel 固定开销会吃掉收益。任务需要 token 间任意长距离交互比如某些全局建模、关系推理任务。模型没有局部注意力先验强行加窗口会牺牲效果。只是想做原型验证不需要关注端到端性能那直接掩码更省事。6.3 回到长期价值控制注意力不是控制“快慢”这一件事FlashAttention 滑动窗口版本的长期价值不只是让 prefill 变快而是让“注意力范围”变成模型设计里一个可调、可解释的维度。以前你想让模型只看窗口只能用 mask 硬遮计算量没省现在 kernel 层面可以跳过无关 block窗口大小才真正成为影响计算资源和效果 trade-off 的旋钮。这意味着你可以更灵活地设计模型局部层和全局层混用、金字塔窗口、动态窗口、甚至窗口大小随层数变化。每一种设计只要最终仍落在“某个 query block 只访问有限 key block”的框架内FlashAttention 的加速机制就依然成立。从工程角度看这才是这套技术最值得关注的地方它不是一次性的性能 hack而是一种更通用的抽象。回到最初的问题FlashAttention 如何加速滑动窗口注意力的 prefill答案是分块加跳过分块解决内存带宽和 softmax 全局归一化的矛盾跳过解决窗口外无效计算的问题。两者合在一起prefill 才真正从 O(N^2) 的负担里解放出来。但每一步都依赖对 block 边界、mask 位置和数值稳定性的严格处理。如果你正在做长文本高效推理可以先从一个小窗口 小序列验证开始把遍历范围画出来再慢慢放大。这样踩坑的成本最低也最容易看清收益到底来自哪里。