这次我们看一个非常具体、但很多人其实没完全想清楚的问题:FlashAttention 到底能不能加速滑动窗口注意力的 prefill 阶段?
先说结论:能,而且理论上比“全量注意力 + FlashAttention”的收益更直接。但前提是 kernel 层面真的做了块级剪枝,而不是只在计算完成后补一个 mask。如果只是把滑动窗口当成一个稀疏 mask 加到普通 FlashAttention 里,那计算量几乎没变,加速也基本谈不上。
这篇文章会把三件事拆开讲清楚:标准注意力慢在哪、FlashAttention 通过什么机制省显存和 IO、滑动窗口注意力如何从 O(n²) 变成 O(n·w),以及它们在 prefill 阶段叠加后到底省了什么。后面会带上可运行的简化代码、理论复杂度推导、工程实现参考和常见误区的排查清单。内容偏原理但落到工程,适合正在做 LLM 推理优化、长上下文部署,或者想搞懂 Mistral 这类滑动窗口模型为什么 prefill 比普通模型快的人。
1. 核心问题速览
先把三个相关概念放在一张表里,明确各自解决什么问题、组合后解决什么问题。
| 概念 | 解决的问题 | 核心思想 | 单独使用的收益 | 与本文主题的关系 |
|---|---|---|---|---|
| 标准 Attention | 无 | 全量 QK^T 后 softmax | 无 | 慢的基线 |
| FlashAttention | 注意力计算中 HBM 访问过多 | 分块 tiling + 在线 softmax + 重计算 | 减少中间矩阵落回显存,显存和 IO 都下降 | 提供高效的“计算内核” |
| 滑动窗口注意力 | 长上下文下注意力矩阵过大 | 每个 token 只关注附近窗口 | 计算量从 O(n²) 降到 O(n·w) | 提供稀疏模式 |
| 两者结合 | prefill 阶段长提示词计算量过大 | 在 FlashAttention 块循环里跳过带外块 | 计算量和 IO 同时下降 | 本文重点 |
从数学角度看,假设序列长度 n = 32768,窗口 w = 2048,双向滑动窗口。朴素注意力的有效注意力对数是 n² ≈ 10.7 亿,滑动窗口注意力是 n·w ≈ 6700 万。也就是说,理论注意力计算量约为原来的 1/16。落入 FlashAttention 的块循环后,如果块大小 B 远小于窗口 w,需要真正计算的块数也大致按这个比例下降。
但注意,这是理论值。真实收益还取决于 kernel 是否真的跳过带外块、块大小和窗口的比率、GPU 调度开销、序列实际长度分布等因素。后面会一一展开。
2. 前置知识:标准 Attention 为什么慢
先看标准注意力公式:
Attention(Q, K, V) = softmax(QK^T / √d) · V如果直接按这个公式实现,一次 prefill 前向需要做这几件事:
- 计算 S = QK^T / √d,得到一个 n×n 的注意力分数矩阵。
- 对 S 做 softmax,按行归一化。
- 用 P 乘 V,得到输出。
问题就出在 n×n 的矩阵上。
当 n = 32768 时,S 矩阵有 10 亿个元素。用 float16 存,就是 2GB。这个矩阵写回显存、再从显存读出来做 softmax、再读出来乘 V,光这一层注意力的 HBM 访问量就是 6GB 以上。而且这是单头单层的量,多头堆叠之后,显存和带宽完全扛不住。
所以标准注意力在长上下文下的瓶颈通常不是 GPU 算力不够,而是 HBM 带宽被中间矩阵的读写占满了。这也是 FlashAttention 出现的核心动机:不要在显存里构造完整 S 矩阵,把计算分块放到 SRAM 里完成。
3. FlashAttention 的核心机制
FlashAttention 做了三件事,分别是分块 tiling、在线 softmax、重计算。
3.1 分块 tiling
把 Q、K、V 都切成大小为 B 的块。对于每个 Q 块,内层循环遍历所有 K 块和 V 块,在 SRAM 中完成这一对小块的分数计算、softmax 累加和加权输出,最终结果再写回 HBM。
这样做的结果就是,完整的 S 矩阵永远不会被构造出来。每个小块的注意力分数是在 SRAM 里算完、用完、丢弃的。
3.2 在线 softmax
朴素 softmax 需要先得到整行分数,找到最大值,再做指数归一化。分块之后,每个 Q 块对应的分数是一批一批进来的,没法先看完整行,所以必须用在线 softmax 方式:
- 维护当前块之前所有分数块的行最大值 m。
- 维护当前累积的指数和 l。
- 遇到新块时,更新最大值,然后按比例 rescale 之前累积的输出。
下面给一个简化的 Python 实现,用来展示在线 softmax 的核心逻辑。这个实现不是实际 CUDA kernel,但能帮助你理解 flash attention 的数值流程:
import numpy as np def flash_attention_tiled(Q, K, V, block_size=64): """ 简化版 FlashAttention 分块实现,不做重计算。 仅用于理解 online softmax 的数值流程。 Q, K, V: shape (n, d) """ n = Q.shape[0] out = np.zeros_like(Q) for i in range(0, n, block_size): q_block = Q[i:i+block_size] m_i = np.full(q_block.shape[0], -np.inf) l_i = np.zeros(q_block.shape[0]) acc = np.zeros_like(q_block) for j in range(0, n, block_size): k_block = K[j:j+block_size] v_block = V[j:j+block_size] # 当前块分数 s_ij = q_block @ k_block.T / np.sqrt(Q.shape[1]) # 更新 running max m_new = np.maximum(m_i, s_ij.max(axis=1)) # 计算当前块的 exp 值 p_ij = np.exp(s_ij - m_new[:, None]) # 对之前累积结果做 rescale alpha = np.exp(m_i - m_new) l_i = alpha * l_i + p_ij.sum(axis=1) # 累积输出 acc = alpha[:, None] * acc + p_ij @ v_block m_i = m_new # 归一化 out[i:i+block_size] = acc / l_i[:, None] return out这段代码的关键是m_new、alpha、l_i三者的更新顺序。它的正确性可以这样验证:当block_size等于整个序列长度时,这个实现退化成标准 softmax attention。实际 kernel 里的 CUDA 实现和这个逻辑一致,只是把外层循环展开到 thread block,把内层循环放到 SRAM 上。
3.3 重计算
反向传播时,需要重新拿到每个位置的注意力分数来计算梯度。FlashAttention 不保存 S 矩阵,而是在反向时重新计算一次前向分数。用额外的计算量换掉了巨大的显存占用。
对于 prefill 阶段,重计算的意义主要在于:即便在训练或推理的梯度计算中,显存也不再随着序列长度二次膨胀。如果是纯推理,重计算不参与前向,影响不大,但 FlashAttention 的 tiling 和 online softmax 机制在推理 prefill 中同样关键。
3.4 FlashAttention 到底省了什么
标准注意力的 HBM 访问量大概是 O(n²),因为 S 矩阵要写一次、读一次,P@V 结果也要写。FlashAttention 的 HBM 访问量从 O(n²) 量级降到一个更小的量级,核心是消除了中间矩阵在 HBM 里的反复读写。
注意:FlashAttention 并不减少 FLOPs。它的优势是让每次从 HBM 读出来的数据被更多次复用,从而把带宽瓶颈大幅缓解。所以 FlashAttention 在 prefill 中通常能带来数倍加速,这个加速本质上是 IO 优化带来的,不是算得更少。
4. 滑动窗口注意力:带状稀疏与 O(n·w)
滑动窗口注意力的想法非常朴素:每个 query 只关注它前后一定范围内的 key。设窗口大小为 w,双向窗口就是:
Attention(i) = softmax( (q_i · K[i-w:i+w]^T) / √d ) · V[i-w:i+w]如果是因果模型(causal),范围就变成[max(0, i-w), i],左侧窗口,右侧不关注。
4.1 带状矩阵视角
在全量注意力矩阵 S 中,滑动窗口对应一个带状稀疏矩阵。只有主对角线附近宽度约 2w 的条带内有值,其余位置都是 -inf(softmax 后会变成 0)。
这种稀疏模式带来的直接收益是:
- 计算量从 O(n²·d) 降到 O(n·w·d)。
- 如果实现得当,KV cache 也可以只保留窗口内的 key/value,显存占用从 O(n) 变成 O(w)。
- prefill 阶段不需要为每一个 query 都计算全部 key 的注意力分数,只需要计算窗口内的。
代表模型包括 Longformer、Mistral 等。需要注意的是,尽管每一层只看局部,但多层堆叠后信息仍然可以在 token 之间间接传递。窗口外信息不是完全丢失,而是通过中间 token 逐步传播。
4.2 窗口不是越大越好
窗口大小决定了“直接注意力范围”。窗口太小,远距离信息需要经过多层传播,模型建模长距离依赖的能力会下降;窗口太大,计算量又回到接近全量注意力的水平。工程上,窗口大小通常和模型层数、序列长度、任务类型一起调。
5. 块级剪枝:FlashAttention 如何适配滑动窗口
这是本文最核心的部分。FlashAttention 本身是一个分块循环。在全量注意力模式下,一个 Q 块需要遍历所有 K 块。滑动窗口模式下,很多 K 块和当前 Q 块完全没有窗口交集,这些块可以直接跳过。
5.1 跳过条件
假设序列长度为 n,块大小为 B,Q 块编号为 qi,KV 块编号为 ki。一个 Q 块覆盖的 token 范围是:
q_start = qi * B q_end = q_start + B - 1一个 KV 块覆盖的 token 范围是:
k_start = ki * B k_end = k_start + B - 1双向窗口大小为 w,当前 Q 块能看到的 key 范围是:
[q_start - w, q_end + w]如果 KV 块范围完全不落在这个区间内,就跳过:
if k_end < q_start - w or k_start > q_end + w: continue5.2 边界块的掩码
不是所有 KV 块都完全落在窗口内。有些块和窗口部分重叠,块内部分行需要被 mask 掉。这里的 mask 规则很简单:
mask[key_col] = 0 if key_col 在窗口内 else -inf注意:mask 必须在计算 softmax 之前加到分数上,并且要参与 online softmax 的 m 和 l 更新。如果直接跳过块,等价于把整个块所有位置的分数都设为 -inf,不影响 m 和 l。但如果块是部分重叠的,必须按行 mask,否则数值会错误。
5.3 简化代码:块剪枝 + 在线 softmax
下面这个实现把第 3 节的 FlashAttention 加上窗口剪枝和边界 mask,用 Python 完整模拟整个流程:
import numpy as np def flash_attention_sliding_window(Q, K, V, window_size, block_size=64): """ 带滑动窗口的 FlashAttention 简化实现。 这里假设双向窗口;因果窗口只需把右侧边界设成 0。 Q, K, V: shape (n, d) """ n = Q.shape[0] d = Q.shape[1] out = np.zeros_like(Q) for qi in range(0, n, block_size): q_start = qi q_end = min(qi + block_size, n) q_block = Q[q_start:q_end] n_q = q_block.shape[0] m_i = np.full(n_q, -np.inf) l_i = np.zeros(n_q) acc = np.zeros_like(q_block) for ki in range(0, n, block_size): k_start = ki k_end = min(ki + block_size, n) # 块级剪枝:判断这个 KV 块是否可能被窗口覆盖 if k_end < q_start - window_size: continue if k_start > q_end + window_size - 1: # 因为是按 ki 递增顺序遍历,这里可以提前 break break k_block = K[k_start:k_end] v_block = V[k_start:k_end] # 计算分数块 s_ij = q_block @ k_block.T / np.sqrt(d) # 构造块内掩码:双向窗口 # s_ij 的行是 query 索引,列是 key 索引 mask = np.zeros_like(s_ij) for i_off in range(n_q): row_q = q_start + i_off for j_off in range(k_end - k_start): col_k = k_start + j_off if abs(row_q - col_k) > window_size: mask[i_off, j_off] = -np.inf s_ij = s_ij + mask # 在线 softmax 更新 m_new = np.maximum(m_i, s_ij.max(axis=1)) p_ij = np.exp(s_ij - m_new[:, None]) alpha = np.exp(m_i - m_new) l_i = alpha * l_i + p_ij.sum(axis=1) acc = alpha[:, None] * acc + p_ij @ v_block m_i = m_new out[q_start:q_end] = acc / l_i[:, None] return out这个实现的正确性在于两点:
- 完全带外块被跳过,或者在遍历过程中提前 break。
- 部分重叠块通过逐位置 mask 处理,保证 softmax 归一化正确。
实际 CUDA kernel 不会用这种逐位置循环写 mask,而是通过 block index 和 row/col 偏移直接计算出有效范围,减少分支开销。
5.4 到底快在哪
全量 FlashAttention 中,一个 Q 块要遍历全部ceil(n / B)个 KV 块。
滑动窗口下,每个 Q 块只需要遍历窗口覆盖的 KV 块,数量大约是:
2w / B + 2所以总计算块数从:
(n/B)²降到:
(n/B) · (2w/B + 2)当 n 远大于 w、窗口远大于块大小时,理论加速比接近:
n / (2w)举个理论示例:n = 32768,w = 2048,B = 128。全量计算的 Q-K 块对数为 256 × 256 = 65536。滑动窗口下,每个 Q 块大约遍历 2×2048/128 + 2 = 34 个 KV 块,总块对数约 256 × 34 = 8704。理论块对数下降约 7.5 倍。
注意:这只是一个理论示例,真实性能还取决于 kernel 调度、GPU 占用率、块大小与窗口边界效应。落到具体硬件上时,需要以实际 profiling 为准,不能只按块对数预测。
5.5 跳过块时,online softmax 为什么不需要额外处理
一个常见疑问是:跳过了带外块,online softmax 维护的 running max 会不会不正确?
不会。因为带外块在正确实现中等价于分数全为 -inf。在线 softmax 中,一个全 -inf 的块对 m_new 没有贡献(max 还是原来的 m),exp 后全是 0,对 l 也没有贡献。所以跳过它和显式计算它,数值上完全等价。
这也是滑动窗口 + FlashAttention 能结合的数学基础:块级剪枝是精确优化,不是近似优化。
6. Prefill 阶段为什么是重点
6.1 Prefill 和 decode 的区别
LLM 推理分成两个阶段:
| 阶段 | 输入 | 计算特点 | 主要瓶颈 |
|---|---|---|---|
| Prefill | 整个 prompt(可能几千 token) | 并行计算所有 token 的 KV 和 logits | 计算量随 n² 增长,受计算和 SRAM 容量限制 |
| Decode | 当前一个 token | 自回归逐 token 生成 | 读取 KV cache,受显存带宽限制 |
Prefill 是计算密集型阶段,因为所有 token 的注意力可以并行算。一个长 prompt 的 prefill 延时,主要取决于注意力层的 FLOPs 和 HBM 访问量。滑动窗口把有效 FLOPs 直接砍到大约 n·w,FlashAttention 又解决掉中间矩阵的 IO 问题,两个优化叠加,prefill 的收益非常明显。
Decode 阶段则不一样。每次只生成一个 token,Q 只有一行。注意力计算量本身是 O(n·d),滑动窗口对它最直接的影响是缩小 KV cache 规模,减少每次要读取的 KV 量。但 decode 阶段更大的影响来自 KV cache 管理和显存带宽,不是 FlashAttention 的 tiling 能单独解决的。
所以如果你重点关心长 prompt 的“首 token 延迟”,滑动窗口 + FlashAttention 的收益非常值得关注。
6.2 理论加速比的边界
上面给的n/(2w)加速比是纯 FLOPs 视角。工程上有几个因素会吃掉一部分理论收益:
- 边界块的开销。如果窗口 w 只比块大小 B 大几倍,边界块占的比例就很高,剪枝收益下降。
- kernel 调度和 wave quantization。GPU 执行 kernel 时有 wave 边界效应,块数不是总能完美打满所有 SM。
- 非注意力层占比。一个 Transformer 层包含 attention、FFN、normalization、embedding。滑动窗口只优化 attention 部分,如果 FFN 占比很高,整体加速比会被稀释。
- 实现复杂度。如果 kernel 分支太多导致 warp divergence 严重,可能比全量 FlashAttention 还慢。
这也是为什么推荐在真实模型上做 profiling,而不是只看复杂度公式。
7. 工程实现参考与框架集成
7.1 flash-attn 库
FlashAttention 官方实现(Dao-AILab/flash-attention)在 flash-attn 2.x 中提供了对窗口注意力的支持。调用时可以直接指定窗口大小,例如:
from flash_attn import flash_attn_func # q, k, v: shape (batch_size, seqlen, nheads, head_dim) # window_size=(left_window, right_window) # 这里左边窗口 2048,右边 0,表示因果滑动窗口 out = flash_attn_func( q, k, v, dropout_p=0.0, causal=True, window_size=(2048, 0) )注意:不是所有版本都有完全一致的参数行为,具体窗口参数语义以你安装的库版本为准。如果你发现flash_attn_func不支持window_size,大概率是版本太老,需要升级到 2.x。
7.2 vLLM 等推理框架
vLLM 等主流推理框架对 Mistral 这类滑动窗口模型有专门支持。由于 vLLM 使用 PagedAttention 管理 KV cache,滑动窗口模型在推理时通常配合“窗口内 KV 保留”策略,超出窗口的 KV 会被淘汰或覆盖。
这一点对生产环境很重要:滑动窗口不只是降低 prefill 计算量,还能显著减少 KV cache 显存占用,让长上下文服务的并发度更高。
7.3 自研 Triton kernel 的参考思路
如果你需要在自定义模型或实验环境中实现滑动窗口 + FlashAttention,可以用 Triton 快速验证。核心思路是:在每个 Q 块的 kernel 内部,根据窗口边界计算出需要遍历的 K 块范围,而不是固定遍历全量。
import triton import triton.language as tl # 伪代码,只展示块循环的窗口范围计算 # q_block_idx 是当前 Q 块索引 # num_kv_blocks 是 KV 块总数 # window_blocks 是窗口折算成块的数量 start_kv = max(0, q_block_idx - window_blocks) end_kv = min(num_kv_blocks, q_block_idx + window_blocks) for kv_idx in range(start_kv, end_kv): # 加载 K/V 块 # 计算分数 # 应用边界 mask # 更新 online softmax 状态 ...用 Triton 的好处是能快速验证剪枝逻辑和数值一致性,缺点是手写满血 kernel 的调度优化空间有限。生产环境优先使用官方库或成熟框架,自己写 kernel 主要用于学习和特殊场景定制。
8. 生产环境中的批量任务与 KV cache 管理
这里对应到实际推理服务中“批量任务怎么处理”的问题。prefill 阶段往往是 decode 的基础,批量请求进入时,prefill 和 decode 会交错调度。
8.1 批量 prefill 时的滑动窗口
批量推理时,每个请求的序列长度可能不同,但窗口大小通常是模型固定的。这就带来一个工程点:kernel 需要处理不同长度的序列,避免把注意力范围外都填充计算。
常见的做法是:
- 按序列长度分组 padding 到接近的块数,减少浪费。
- 对不同请求使用统一的窗口参数,kernel 内部根据实际长度裁剪遍历范围。
- 结合 continuous batching,让新的 prefill 请求插到 decode 间隙里执行,提高 GPU 利用率。
8.2 KV cache 何时淘汰
滑动窗口模型的 KV cache 不需要保存全部历史 key/value。工程上常见的做法是:
- 超过窗口的旧 KV 直接丢弃。
- 如果使用 PagedAttention 这类块级管理,淘汰粒度为 block,可能存在窗口边界和 block 边界不一致的问题,需要额外处理。
注意:如果你只是用了 FlashAttention 的window_size参数做注意力计算,但 KV cache 没有同步淘汰,那显存收益会大打折扣。计算加速和显存淘汰是两个层面的优化,要一起配齐。
8.3 输出验证
生产环境接入滑动窗口前,建议做一次数值验证:
- 写一个朴素的 version(可以小规模跑 512 token 以下)。
- 用 FlashAttention + 滑动窗口的 version 跑同样输入。
- 对比两个版本输出的 logits 差。
最大绝对误差通常应该在 1e-3 量级以内。如果误差很大,优先检查 mask 是否正确加到了 softmax 之前,以及在线 softmax 的 m/l 更新顺序是否正确。
9. 性能观察与调优实验设计
9.1 怎么观察收益
不要只盯着 end-to-end 总延迟。用 profiling 工具把 attention kernel 的时间拆出来看。推荐观察这几个指标:
- 单个注意力 kernel 的时间。
- kernel 内部的 HBM 吞吐。
- 有效计算块数和实际访问的 KV 块数。
- prefill 阶段峰值显存。
如果你用 PyTorch,可以用torch.profiler先拿到粗粒度数据。想进一步看硬件指标,用 Nsight Compute 或ncu看 attention kernel 的 memory throughput 和 compute throughput。
9.2 一组值得做的对比实验
建议按下面的矩阵做控制变量实验:
| 变量 | 实验组 |
|---|---|
| 序列长度 | 4096、8192、16384、32768 |
| 窗口大小 | 512、1024、2048、全量 |
| 注意力实现 | 朴素 mask 滑动窗口、FlashAttention 全量、FlashAttention + 滑动窗口 |
| 批量大小 | 1、2、4、8 |
记录结果时,至少包含:
- prefill 延迟。
- attention kernel 时间占比。
- 峰值显存。
- 数值一致性检查通过与否。
这种实验设计能帮你判断:在你的硬件和模型上,滑动窗口 + FlashAttention 的收益主要来自哪个维度。
9.3 调优思路
如果发现收益不如预期,按下面的顺序排查:
- 确认 attention kernel 真的跳过了带外块,而不是只做了 mask。用 profiler 看 kernel 内循环的 block 访问范围。
- 确认块大小设置。块太大,边界块占比高;块太小,调度开销大。通常 B 取 64 或 128,具体要看 GPU 架构。
- 确认瓶颈不在 attention 之外。如果 FFN 占比过高,滑动窗口优化对总延迟的贡献有限。
- 确认没有因为窗口淘汰 KV cache 导致上下文丢失过多,从而影响输出质量。
10. 常见误区和排查方法
| 误区或问题 | 原因 | 排查方式 | 正确理解或解决方案 |
|---|---|---|---|
| 滑动窗口 + FlashAttention 只是补一个 mask | 没有在块循环层面跳过带外块 | 检查 kernel 里 Q 块遍历 K 块的索引范围 | 必须做块级剪枝,单纯 mask 不会省计算量 |
| 用了窗口参数但速度没变化 | kernel 没真正生效,或序列太短 | profiling 看 attention kernel 时间 | 序列要远大于窗口,收益才能体现 |
| 输出和朴素窗口注意力不一致 | mask 位置错误、online softmax 的 rescale 顺序错误 | 对比 512 token 下朴素实现和优化实现 | 检查 mask 是否加在 exp 之前 |
| 显存没有下降 | 只优化了注意力计算,KV cache 没有淘汰 | 观察 KV cache 显存占用 | 需要配合滑动窗口的 KV cache 管理 |
| 窗口设小后效果变差 | 任务确实需要全局依赖 | 对比不同窗口的输出困惑度或下游指标 | 考虑多层局部注意力叠加是否足够,或者改用全局 token |
| 长提示词 prefill 仍然 OOM | attention 之外的层占用峰值显存 | 逐步 profiling 每层显存 | 关注 embedding、FFN、KV cache 的峰值 |
| 跳过块后 running max 错误 | 担心全 -inf 块影响在线 softmax | 数值一致性测试 | 理论上跳过全 -inf 块不会影响 m/l,因为不改变 max 和 sum |
10.1 关于“窗口外信息是否完全丢失”
这个需要说清楚。滑动窗口确实把窗口外的分数 mask 成了 -inf,当前层计算时完全不看它们,但这不代表最终输出里窗口外信息完全无法影响某个位置的表示。Transformer 多层堆叠后,信息可以通过邻居 token 逐步“接力”传播。所以窗口大小和层数共同决定了有效感受野。
如果任务本身需要很强的长距离直接依赖(比如长文档的全局指代消解),窄窗口可能不够。如果任务以局部语义为主,窄窗口加上 FlashAttention 的提速非常划算。
11. 最佳实践与总结
11.1 最佳实践清单
- 先验证数值一致性,再优化性能。把朴素滑动窗口注意力作为 reference,和 flash attention + sliding window 对比。
- 窗口和块大小的关系要清楚。块大小远小于窗口时剪枝效率最高;窗口很小时边界块占比高,收益可能下降。
- prefill 和 decode 分开评估。滑动窗口对 prefill 的计算收益和对 decode 的 KV cache 收益不是一回事。
- 生产环境优先使用成熟框架。flash-attn 的 window_size 参数、vLLM 对滑动窗口模型的支持,都比自己写 kernel 稳定。
- KV cache 淘汰和计算剪枝要配套。只做计算剪枝不做缓存淘汰,显存收益会被浪费。
- 任务评估不能只看延迟。窗口越小,推理越快,但输出质量可能变差。需要在下游任务上做质量对比。
- 长序列推理时预留一定显存余量。即使理论计算量降下来了,峰值显存还受 batch、FFN、输入长度共同影响。
11.2 回到最初的问题
FlashAttention 能不能加速滑动窗口注意力的 prefill?
能。
加速来自两个层面:
- FlashAttention 消除中间注意力矩阵的 HBM 读写,让注意力计算本身更快。
- 滑动窗口让 FlashAttention 的块循环可以跳过大量无窗口交集的 KV 块,让有效计算量从 O(n²) 降到 O(n·w)。
两者组合之后,prefill 阶段的长提示词处理能力会有量级上的改善。但前提是 kernel 真正实现了块级剪枝,并且窗口、块大小、序列长度之间的比例关系选得合理。
建议收藏备用。如果你正在研究长上下文推理优化,或者准备在自己的模型里接入滑动窗口,建议先从文中第 5 节的简化代码入手跑通数值验证,再切换到官方库或框架实现。最容易踩的坑就是“以为加了 mask 就等于做了滑动窗口优化”,先绕开它,后面的路会顺很多。