如果正在做长文本大模型的推理优化,大概率会遇到这样一个场景:prompt 已经到几千 token,prefill 阶段 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 从头到尾都不需要完整存在。
可以这样对比:
| 对比维度 | 朴素 Attention | FlashAttention |
|---|---|---|
| 中间矩阵 S/P | 需要完整写回 HBM | 在 SRAM 中局部计算,不写回 |
| 计算调度 | 矩阵乘、softmax、矩阵乘分开执行 | 融合到分块循环中 |
| softmax 处理 | 整行先算 max 和 sum | online 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 的思路加速窗口 prefill
5.1 窗口对齐到块:最关键的一步
FlashAttention 的分块逻辑中,一个 block 是大小为 B×B 的矩阵块。滑动窗口的有效范围是斜带状,如果要让一个 block 要么完全在窗口内、要么完全在窗口外,最稳妥的做法是把窗口大小 W 设置成 block size B 的整数倍。
假设 B = 32,W = 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_size=32, window_size=64): """ 教学用分块实现:窗口 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_size=64) out_flash = flash_window_attention(q, k, v, block_size=32, window_size=64) 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 生成时都要读取全部历史 KV,decode 阶段每步的访存成本随上下文长度线性上升。窗口注意力提出了一种更经济的方案:位置 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 + mask | O(n²) 实际仍全算 | n×n mask 和中间矩阵 | 否 |
| 稀疏 mask + FlashAttention 库 | O(nW) 但未做 block 对齐 | 依赖 mask 实现,可能仍有额外开销 | 部分 |
| 窗口 block 级跳过 + 融合 kernel | O(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 或 Inf | mask 中使用了-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 级跳过;