ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

FlashAttention加速滑动窗口注意力Prefill的工程解析

FlashAttention加速滑动窗口注意力Prefill的工程解析 这次我们看一个非常具体、但很多人其实没完全想清楚的问题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_size64): 简化版 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:iblock_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:jblock_size] v_block V[j:jblock_size] # 当前块分数 s_ij q_block k_block.T / np.sqrt(Q.shape[1]) # 更新 running max m_new np.maximum(m_i, s_ij.max(axis1)) # 计算当前块的 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(axis1) # 累积输出 acc alpha[:, None] * acc p_ij v_block m_i m_new # 归一化 out[i:iblock_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 矩阵要写一次、读一次PV 结果也要写。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:iw]^T) / √d ) · V[i-w:iw]如果是因果模型causal范围就变成[max(0, i-w), i]左侧窗口右侧不关注。4.1 带状矩阵视角在全量注意力矩阵 S 中滑动窗口对应一个带状稀疏矩阵。只有主对角线附近宽度约 2w 的条带内有值其余位置都是 -infsoftmax 后会变成 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块大小为 BQ 块编号为 qiKV 块编号为 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_size64): 带滑动窗口的 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(axis1)) p_ij np.exp(s_ij - m_new[:, None]) alpha np.exp(m_i - m_new) l_i alpha * l_i p_ij.sum(axis1) 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 32768w 2048B 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 还是原来的 mexp 后全是 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·wFlashAttention 又解决掉中间矩阵的 IO 问题两个优化叠加prefill 的收益非常明显。Decode 阶段则不一样。每次只生成一个 tokenQ 只有一行。注意力计算量本身是 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_p0.0, causalTrue, 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 仍然 OOMattention 之外的层占用峰值显存逐步 profiling 每层显存关注 embedding、FFN、KV cache 的峰值跳过块后 running max 错误担心全 -inf 块影响在线 softmax数值一致性测试理论上跳过全 -inf 块不会影响 m/l因为不改变 max 和 sum10.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 就等于做了滑动窗口优化”先绕开它后面的路会顺很多。
返回列表