先提醒一句:Attention 机制我们天天在用,但真正决定大模型推理时“首字延迟”的,往往不是 decode,而是 prefill。而 prefill 阶段最有性价比的优化组合之一,就是 FlashAttention + 滑动窗口注意力。本文从原理到伪代码,完整拆解这套组合为什么能加速、能快在哪、落地时要注意什么。适合正在做推理优化、长文本模型部署,或者想深入理解 attention kernel 的读者。
1. 背景与核心概念
1.1 先搞清楚 prefill 是什么
大模型生成回答时,分成两个阶段:
- prefill(预填充)阶段:用户输入整段 prompt,模型一次性并行处理所有 token,生成第一个输出 token 前的计算阶段。
- decode(解码)阶段:每步只生成一个 token,需要反复读取历史 token 的 KV cache。
prefill 阶段的特点是:输入序列很长、计算量极大、并行度极高。对于用户来说,prefill 耗时直接决定“首字延迟”,也就是屏幕上第一个字出现的速度。一个 4k token 的输入,在未优化场景下可能让 prefill 花费数百毫秒甚至几秒。
很多做推理优化的朋友,一开始把目光都放在 decode 的 KV cache 上,反而忽略了 prefill。但对于长文档问答、代码补全、Agent 场景的 long context 输入,prefill 往往是系统瓶颈。
1.2 滑动窗口注意力解决什么问题
标准 Transformer 的注意力是全局的:每个 query 要和所有 key 计算相关度。这带来了两个问题:
- 计算复杂度是 O(N²),N 是序列长度。
- 显存占用同样是 O(N²),因为要保存 N×N 的注意力矩阵。
但直觉上,一个 token 往往只和附近若干 token 有较强关联。滑动窗口注意力(Sliding Window Attention,SWA)正是基于这一假设:每个 query 只允许关注窗口 W 范围内的 key。
以 Mistral 系列模型为代表的许多开源模型,都采用了滑动窗口注意力作为长文本场景的稀疏化策略。窗口注意力把标准注意力的“全连接”稀疏化成“带状”结构。
1.3 FlashAttention 是什么
FlashAttention 是一种 IO 感知(IO-aware)的精确注意力算法。它解决的核心问题不是减少计算量,而是减少显存访问开销。
传统注意力实现流程是:
- 计算 S = Q @ K^T,得到 N×N 的注意力分数矩阵。
- 对 S 逐行做 softmax,得到 P。
- 计算 O = P @ V。
问题在于 S 和 P 都是 N×N 的中间矩阵,需要写入高带宽内存(HBM),随后又被读回。Attention 的计算量是 O(N²),中间矩阵的显存访问量也是 O(N²)。在长序列场景下,访存开销甚至超过计算开销。
FlashAttention 通过分块 tiling、online softmax、kernel 融合等技术,把中间矩阵留在芯片上的 SRAM(片上存储)中,避免频繁读写 HBM。它并没有改变注意力计算的结果,是精确算法,而非近似注意力。
1.4 为什么要把这三者放在一起讨论
在实际推理系统中,滑动窗口注意力和 FlashAttention 经常同时出现:
- 滑动窗口注意力负责“减少计算量”,通过稀疏化去掉不重要的 key。
- FlashAttention 负责“减少访存量”,让真正要算的部分更高效地完成。
尤其在 prefill 阶段,序列长、计算密集,两者结合能解决“计算量大”和“访存开销大”两方面的双重瓶颈。这也是本文最想讲透的地方。
2. FlashAttention 的加速原理拆解
2.1 从算力到带宽:Attention 的性能瓶颈转移
先看一个基本事实:现代 GPU 上,一个 kernel 的性能可能受两种因素限制:
- Compute-bound(计算受限):矩阵乘法、卷积这类算子,计算量远大于数据搬运量,瓶颈是 GPU 的浮点算力。
- Memory-bound(访存受限):元素级操作、softmax、LayerNorm 这类算子,计算很简单,瓶颈是从 HBM 搬运数据的速度。
标准注意力实现里,QK^T 是计算受限的 GEMM,但紧随其后的 mask + softmax + dropout + P@V 中,S 和 P 矩阵需要先写回 HBM 再读回。这部分访存开销很容易压过计算收益。
可以用一个简单对比来理解:
| 阶段 | 标准实现 | FlashAttention |
|---|---|---|
| 注意力分数矩阵 S | 写入 HBM | 留在 SRAM |
| 概率矩阵 P | 写入 HBM | 留在 SRAM |
| softmax 归一化 | 需要读 S 两次 | online 方式一次扫描 |
| 中间矩阵显存 | O(N²) | O(N) |
所以 FlashAttention 的核心思想可以概括为:不要让中间结果去“长途旅行”,在片上 SRAM 里完成尽可能多的计算。
2.2 Tiling:把大矩阵切成能装进 SRAM 的小块
SRAM 的特点是速度快但容量小。以 A100 为例,每个 SM 上的 SRAM(shared memory)大约 100~200KB 级别,远装不下 N×N 的注意力矩阵,但装得下一个 N×N 的子块。
FlashAttention 的做法是:
- 把 Q 矩阵按行切成若干个 block,记为 Q_i。
- 把 K、V 矩阵也按行切块,记为 K_j、V_j。
- 每次只加载一个 Q_i 和一组 K_j、V_j 到 SRAM。
- 在 SRAM 中完成 Q_i @ K_j^T,得到局部的注意力分数。
- 用 online softmax 更新输出 O_i。
这个过程相当于把大矩阵乘法拆成了“在外层循环里遍历 K/V 块,在内层将结果累积到输出块”的形式。由于所有中间分数块都在 SRAM 中,HBM 传输量大幅下降。
2.3 Online Softmax:不用保存完整分数矩阵就能归一化
传统 softmax 需要先遍历一整行,找到最大值,再算指数和,最后归一化。这意味着至少要把注意力分数矩阵完整“看过一遍”才能继续。
FlashAttention 使用 online softmax 的技巧:在遍历 K/V 块时,实时维护当前行的局部最大值 m_i 和局部指数和 l_i。
每处理一个新的 block,就更新:
- 新最大值 m_new = max(m_old, 当前块最大值)
- 修正因子 = exp(m_old - m_new)
- 指数和 l_new = l_old * 修正因子 + 当前块指数和
- 输出 O 也要按比例修正
这样只需要一遍扫描就能得到正确的 softmax 结果。最终输出是 O / l,和标准 softmax 计算出的结果在数学上完全等价。
2.4 重计算:用一次 extra 的矩阵乘法换显存
反向传播时,标准注意力需要保存 S 和 P 矩阵用于梯度计算。FlashAttention 选择不保存,而是在反向传播时重新计算一遍正向的注意力分数。多了一次计算,但省掉了 O(N²) 的显存占用。
这在训练场景回非常关键,因为显存直接决定了最大的 batch size 和序列长度。在推理场景,正向重计算收益一般,但 prefill 阶段如果涉及输入敏感的服务,显存省下来就能支持更长的上下文。
2.5 小结
FlashAttention 的三板斧:
- tiling:适配 SRAM 容量,避免大中间矩阵落回 HBM。
- online softmax:精确计算,一次扫描完成归一化。
- kernel 融合:把 QK^T、mask、softmax、P@V 融合成一个 kernel,减少 kernel 启动和全局内存读写。
理解了这三个点,就能理解为什么 FlashAttention 对长序列的加速特别明显:它把“计算和访存的比值”重新拉回了有利于 GPU 算力发挥的区域。
3. 滑动窗口注意力与 FlashAttention 的结合点
3.1 滑动窗口注意力的稀疏结构
滑动窗口注意力中,每个 query 的注意力范围是:
[key_index ∈ [query_index - W + 1, query_index]]在注意力矩阵中,每一行只有连续的 W 个位置是非零的。整体形成一条带状矩阵(band matrix)。
对于 causal + sliding window 的设置,每个 query 只能看到不超过 W 个历史 token。窗口 W 通常是 512、1024、2048 这样的值。
3.2 直接掩码的问题:计算没省,显存照样爆
很多初学者会在 PyTorch 里这样实现滑动窗口注意力:
mask = torch.tril(torch.ones(N, N)) # 再叠加窗口 window_mask = torch.tril(torch.ones(N, N), diagonal=0) - torch.tril(torch.ones(N, N), diagonal=-W)然后把这个 mask 加到 attention score 上。
这段代码有两个问题:
- mask 是 N×N 的稠密矩阵,即使它表示的是稀疏结构,显存占用仍然是 O(N²)。
- Q @ K^T 仍然全量计算,只是把不需要的位置用负无穷遮掉,FLOPs 一点没省。
也就是说,这种写法虽然“语义上”是滑动窗口注意力,但性能和标准注意力没有任何区别。
真正省计算的实现必须做到“计算时不访问窗口外的 key”。
3.3 FlashAttention 的 tiling 天然适配窗口结构
FlashAttention 的遍历单位是 block,不是单个 token。假设 block_size = 64,查询块的索引范围是 [q_start, q_start + 64],那么对应的有效 key 范围是:
[q_start - W + 1, q_start + 63]这意味着,在遍历 K/V 块时,不需要遍历整个序列,只需要遍历落在上述区间内的块即可。
相比稠密注意力需要遍历 N 个 key 块,滑动窗口只需要遍历约 (W + block_size) / block_size 个 key 块。这个减少是线性的:如果 N = 8192,W = 1024,那么遍历的 key 块数量从 128 块降到了约 17 块。
注意:这里减少的是“遍历路径”,而不是完全严格的 W block。因为块边界的存在,实际遍历范围比严格窗口略大。
3.4 块边界不对齐问题
假设 block_size = 64,窗口 W = 1024。窗口大小刚好是 block_size 的 16 倍,边界对齐很完美。
但如果 W = 1000,窗口边界就不会和块边界对齐。实际实现通常有两种处理方式:
- 向上取整到 block 对齐:把 W 对齐到 block_size 的整数倍,多算几个 key。这种做法实现简单,只有极少量的额外计算。
- 在块内做细粒度 mask:加载一个 key block 后,在块内逐 token 判断是否在窗口内,不在的位置用 -inf 填充。
FlashAttention 的风格倾向于第二种:块之间用“跳过”,块内部用 mask。这样既能保证计算量接近理论值,又能避免 padding 带来的浪费。
3.5 结合后的计算流程
最终,结合后的 prefill 注意力 kernel 流程是这样的:
- 外层循环遍历 query block。
- 根据当前 query block 的起始位置,计算需要遍历的 key block 区间。
- 跳过落在窗口外的 key block。
- 对窗口内的 key block,加载到 SRAM。
- 计算局部 QK^T,块内应用 mask(causal + 窗口边界)。
- 用 online softmax 累加结果。
- 写回输出。
这套流程同时做到了两件事:
- 计算量从 O(N²) 降到约 O(N × W)。
- 访存量也从 O(N²) 级别的中间矩阵降到 O(N × W) 级别的 block 数据加载。
4. 为什么 prefill 阶段加速收益最大
4.1 prefill 的计算量分析
先拆解 prefill 阶段 attention 部分的计算量。
标准注意力的 FLOPs:
- QK^T:2 × N × N × d_head
- P @ V:2 × N × N × d_head
- 合计约 4 × N² × d_head
滑动窗口注意力的 FLOPs:
- QK^T:2 × N × W × d_head
- P @ V:2 × N × W × d_head
- 合计约 4 × N × W × d_head
当 N 远大于 W 时,计算量的下降接近 W/N 的比例。假设输入长度 32k,窗口 1024,attention 部分的计算量可以降到原来的约 1/32。
但这里必须说明:attention 只占 prefill 全部计算量的一部分。QKV 投影(三个线性层)的计算量是 3 × N × d_model²,这部分在全连接层占比越高,attention 稀疏化的收益就越被稀释。通常 d_model 不大或序列极长时,attention 才是主导。
4.2 prefill 的访存量分析
prefill 阶段,Q、K、V 都是完整的序列级矩阵,需要一次性从 HBM 读取。但最要命的是中间注意力矩阵 S 和 P。
标准实现的 S 矩阵大小是 N² × 2 bytes(如果 fp16),对 32k 序列就是 32k × 32k × 2 ≈ 2GB。这还只是单层单头。
FlashAttention 解决了中间矩阵的访存问题,而滑动窗口进一步减少了需要从 HBM 加载的 K/V 数据量。两者叠加后,prefill 的访存开销从“O(N²) 中间矩阵 + O(N²) K/V 读取”下降为“O(N×W) 的 K/V 读取”。
4.3 decode 阶段的收益为什么有限
decode 阶段每个 step 只有一个 query token,Q 的形状是 [1, d_model],K/V 从 KV cache 中读取。
这时候:
- 注意力分数矩阵是 [1, N],非常小。
- 真正的大头是读取 KV cache 的带宽。
滑动窗口注意力在这个阶段的收益在于:KV cache 可以裁剪到窗口大小,显存占用降低。但如果系统已经使用了 KV cache 的 LRU 淘汰、chunked prefill 等优化,滑动窗口能带来的边际收益没有 prefill 那么明显。
换句话说:
- prefill 是计算密集 + 中间矩阵访存密集,FlashAttention 和稀疏化都能发挥最大作用。
- decode 是 KV cache 访存密集,主要靠 KV cache 管理和并行策略优化。
5. 伪代码级实现示意
这一节进入核心实现思路。需要说明,这里的代码是教学示意,用于表达 kernel 的遍历逻辑和窗口处理方式,不是可以直接运行的性能最优代码。
5.1 基础数据结构
# 示意:以 Python 类描述 FlashAttention + Sliding Window 的 kernel 配置 @dataclass class WindowAttentionConfig: block_size_q: int = 64 # query block 大小 block_size_k: int = 64 # key block 大小 window_size: int = 1024 # 滑动窗口大小 causal: bool = True # 是否 causal mask num_heads: int = 32 head_dim: int = 128窗口大小为 W,表示当前 query 最多能看到包括自身在内的前 W 个 key。
5.2 窗口范围内 key block 的索引计算
def get_key_block_range(q_block_start, total_seq_len, window_size, block_size_k): """ 给定 query block 的起始 token 位置,计算需要遍历的 key block 区间。 """ # 当前 query block 的 token 范围是 [q_block_start, q_block_start + block_size_q - 1] # 最早需要看到的 key 位置 key_start = max(0, q_block_start - window_size + 1) # 最晚需要看到的 key 位置(如果是 causal,就是当前 block 的最后一行) key_end = q_block_start + block_size_q # 开区间 # 换算成 key block 的索引 block_start = key_start // block_size_k block_end = (key_end + block_size_k - 1) // block_size_k return block_start, block_end这个函数的目的是让外层循环跳过完全落在窗口外的 key block。
5.3 单 query block 的遍历流程
def flash_attention_windowed(Q_i, K, V, config): """ 处理一个 query block 的简化流程。 Q_i: [block_size_q, head_dim],当前 query 块 K, V: [seq_len, head_dim],完整序列的 key/value 返回: O_i [block_size_q, head_dim] """ block_size_q = config.block_size_q block_size_k = config.block_size_k W = config.window_size seq_len = K.shape[0] # 当前 query block 的起始位置 q_start = 当前遍历的 query 起始位置 # online softmax 状态 m_i = torch.full((block_size_q,), -float('inf'), device=Q_i.device) l_i = torch.zeros((block_size_q,), device=Q_i.device) O_i = torch.zeros_like(Q_i) # 计算需要遍历的 key block 区间 k_start_block, k_end_block = get_key_block_range( q_start, seq_len, W, block_size_k ) for j in range(k_start_block, k_end_block): K_j = K[j * block_size_k : (j + 1) * block_size_k, :] V_j = V[j * block_size_k : (j + 1) * block_size_k, :] # 1. 计算当前块的注意力分数 S_ij = Q_i @ K_j.T # [block_size_q, block_size_k] # 2. 在块内生成 mask mask = construct_window_mask(q_start, j * block_size_k, block_size_q, block_size_k, W, config.causal) S_ij = S_ij.masked_fill(mask, -float('inf')) # 3. online softmax 更新 m_ij = S_ij.max(dim=-1, keepdim=True).values m_new = torch.maximum(m_i, m_ij.squeeze(-1)) # 修正比例 alpha = torch.exp(m_i - m_new) beta = torch.exp(m_ij.squeeze(-1) - m_new) # 4. 更新输出 P_ij = torch.exp(S_ij - m_new.unsqueeze(-1)) O_i = O_i * alpha.unsqueeze(-1) + P_ij @ V_j # 5. 更新统计量 l_i = l_i * alpha + P_ij.sum(dim=-1) m_i = m_new # 最终归一化 O_i = O_i / l_i.unsqueeze(-1) return O_i5.4 窗口 mask 的构造
def construct_window_mask( q_start, k_block_start, block_size_q, block_size_k, window_size, causal ): """ 构造当前 block 的 mask。 返回 [block_size_q, block_size_k] 的 bool 矩阵,True 表示需要屏蔽。 """ q_idx = torch.arange(q_start, q_start + block_size_q).unsqueeze(1) k_idx = torch.arange(k_block_start, k_block_start + block_size_k).unsqueeze(0) # 1. causal mask causal_mask = k_idx > q_idx # 2. 窗口 mask window_mask = k_idx < (q_idx - window_size + 1) mask = causal_mask | window_mask return mask注意这里q_idx和k_idx都是相对的,实际 kernel 中需要传入真实的 token 位置,否则在有 padding 或使用 varlen 输入时会出错。
5.5 外层循环的调用
def flash_attention_windowed_full(Q, K, V, config): """ 完整序列的 prefill 阶段。 Q, K, V: [seq_len, head_dim] 返回 O: [seq_len, head_dim] """ seq_len = Q.shape[0] O = torch.zeros_like(Q) for i in range(0, seq_len, config.block_size_q): Q_i = Q[i : i + config.block_size_q, :] O_i = flash_attention_windowed(Q_i, K, V, config) O[i : i + config.block_size_q, :] = O_i return O这个外层循环就是 FlashAttention 的“query block 遍历”,而内层j循环则被限制在窗口范围内。
5.6 从伪代码到真正的 CUDA kernel
真正的实现中,Q_i、K_j、V_j被显式加载到 shared memory:
- 每个
K_j、V_j从 HBM 拷贝到 SRAM。 S_ij的中间矩阵只存在于寄存器中。O_i一直留在寄存器中,只有最终结果写回 HBM。
这是 FlashAttention 的核心优化,也是伪代码与真实实现差距最大的地方。理解思路后,建议直接阅读现有开源实现,比如 Triton 版的 FlashAttention 教程或 flash-attn 库的源码。
6. 工程实践:性能与实现的几个关键点
6.1 Block 大小与窗口的协调
block_size 的选择会影响实际性能。假如 block_size = 64,窗口 W = 1000,那么 key block 遍历范围会覆盖 [972, 1063] 之类的区间,实际计算的 key 数比理论窗口略多。
选择 block 时可以考虑:
- 让窗口尽量是 block_size 的整数倍,减少块内无效计算。
- block_size 不宜过大,因为 SRAM 容量有限,Q_i、K_j、V_j 都占 shared memory。
- block_size 不宜过小,否则循环次数多,kernel 启动和索引计算开销占比上升。
正常实践中,block_size 常在 64 到 128 之间。
6.2 与 GQA / MQA 的配合
很多模型的 attention 用的是 GQA(分组查询注意力)或 MQA(多查询注意力),其中多个 query head 共享一组 KV head。
这意味着在加载 K/V block 时,只需要加载一次,多个 query head 复用。FlashAttention 的 tile 设计本身就支持这种复用,滑动窗口的遍历范围是 head 无关的,所以两者可以自然叠加。
工程实现时需要注意 shared memory 的分配策略:K/V block 按 group 加载,Q block 按 head 依次处理。
6.3 长序列下的 varlen 处理
真实场景中,一个 batch 内可能有多个样本,长度各不相同。每个请求的窗口遍历范围不同,不能简单用固定seq_len做索引。
开源库通常提供 varlen 接口,传入每个样本的起始位置(cu_seqlens)和窗口参数。实现时需要在 kernel 内部根据q_start判断当前样本边界,防止跨样本计算注意力。
6.4 profiling 时观察什么指标
优化完成后,需要验证确实把访存降下来了。用 Nsight Compute 等工具观察:
- SM 占用率:是否打满。
- shared memory 使用量:是否还在合理范围。
- HBM 读写量:这是核心指标,FlashAttention + SWA 应该显著低于标准实现。
- kernel 耗时:对比稠密 FlashAttention,滑动窗口版本的耗时应该随 W/N 的比例下降。
如果只看到计算量下降但 HBM 读写量没变,说明窗口外的 key block 没有真正跳过,可能只是 mask 掉了。
6.5 什么时候不该用滑动窗口
- 需要精确忽略远处 token 的任务:有些任务(如全文总结、跨章节推理)需要全局注意力,滑动窗口会牺牲召回能力。
- 序列长度不长的场景:N = 512 时,N² 与 N×W 的差距不明显,不值得引入稀疏化复杂度。
- 已经用其他稀疏策略:如果模型本身是 Longformer 式的 dilated 滑动窗口,或者有全局 token 特化设计,需要综合考虑。
7. 常见误区与排查思路
| 误区 | 实际情况 | 正确理解 |
|---|---|---|
| FlashAttention 是近似注意力 | FlashAttention 是精确算法 | 它不改变注意力结果,只改变计算的访存模式 |
| 稀疏注意力一定比稠密快 | 在短序列时不一定 | 需要 N 远大于 W 且 kernel 真正跳过窗口外 block 才有效 |
| 实现 SWA 只需要加 mask | 仅加 mask 不省计算和显存 | 需要修改计算遍历路径,才能获得收益 |
| prefill 优化没用,反正 decode 是瓶颈 | 长上下文输入时 prefill 是首字延迟瓶颈 | 滑动窗口在 prefill 阶段收益最大 |
| 窗口越小越好 | 窗口过小会导致模型能力明显下降 | 需要在质量与性能之间取舍 |
一个常见的排查场景是:模型开启滑动窗口后,prefill 速度没变化。
排查步骤:
- 检查是否真的跳过了窗口外的 key block。
- 检查窗口是否被 mask 正确应用在 block 内部。
- 检查序列长度 N 是否远大于窗口 W。
- 检查是否 batch 内 padding 导致遍历仍然覆盖全序列。
- 检查 QKV 投影层是否成为新的瓶颈,attention 时间占比是否已经很小。
8. 总结与学习路线
FlashAttention 与滑动窗口注意力的结合,本质上是两件事的叠加:
- FlashAttention 让 attention kernel 更贴近 SRAM 的读写极限,减少 HBM 访存。
- 滑动窗口让 attention 真正跳过不需要计算的位置,减少计算量和加载量。
两者在 prefill 阶段能产生“1 + 1 > 2”的效果:因为 prefill 既需要处理完整的 Q/K/V 序列,又需要处理 N×N 级的大矩阵,这两个优化方向正好命中了 prefill 的两大痛点。
如果你想继续往这个方向深入,建议按下面的路径走:
- 先手写一个标准 attention 的 PyTorch 实现,用 profiling 工具确认访存瓶颈。
- 用 Triton 实现一个简化版 FlashAttention,跑通 2k 序列长度。
- 在简化版基础上增加窗口跳跃逻辑,对比不同 W 下的耗时曲线。
- 阅读官方 flash-attn 库中 window_size 相关实现,理解 cu_seqlens 和 block table 的设计。
- 在推理框架中接入已经封装好的滑动窗口注意力 kernel,观察真实业务场景的 prefill 延迟变化。
动手时可以先做一个小实验:在 Triton 里实现一个窗口大小为 512、序列长度为 8192 的 attention kernel,和稠密 FlashAttention 对比 prefill 耗时。当你能解释清楚“为什么耗时不是严格的 16 倍下降,而是更接近 6~8 倍”这个问题时,你就真正理解了这个主题。