
先提醒一句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 AttentionSWA正是基于这一假设每个 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 PV 中S 和 P 矩阵需要先写回 HBM 再读回。这部分访存开销很容易压过计算收益。可以用一个简单对比来理解阶段标准实现FlashAttention注意力分数矩阵 S写入 HBM留在 SRAM概率矩阵 P写入 HBM留在 SRAMsoftmax 归一化需要读 S 两次online 方式一次扫描中间矩阵显存O(N²)O(N)所以 FlashAttention 的核心思想可以概括为不要让中间结果去“长途旅行”在片上 SRAM 里完成尽可能多的计算。2.2 Tiling把大矩阵切成能装进 SRAM 的小块SRAM 的特点是速度快但容量小。以 A100 为例每个 SM 上的 SRAMshared 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、PV 融合成一个 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), diagonal0) - 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 8192W 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块内应用 maskcausal 窗口边界。用 online softmax 累加结果。写回输出。这套流程同时做到了两件事计算量从 O(N²) 降到约 O(N × W)。访存量也从 O(N²) 级别的中间矩阵降到 O(N × W) 级别的 block 数据加载。4. 为什么 prefill 阶段加速收益最大4.1 prefill 的计算量分析先拆解 prefill 阶段 attention 部分的计算量。标准注意力的 FLOPsQK^T2 × N × N × d_headP V2 × N × N × d_head合计约 4 × N² × d_head滑动窗口注意力的 FLOPsQK^T2 × N × W × d_headP V2 × N × W × d_head合计约 4 × N × W × d_head当 N 远大于 W 时计算量的下降接近 W/N 的比例。假设输入长度 32k窗口 1024attention 部分的计算量可以降到原来的约 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 tokenQ 的形状是 [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), deviceQ_i.device) l_i torch.zeros((block_size_q,), deviceQ_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, keepdimTrue).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 倍”这个问题时你就真正理解了这个主题。