公司动态
FlashAttention加速滑动窗口注意力的prefill优化
在优化大模型推理性能时prefill 阶段往往是最容易暴露问题的地方。尤其是当 prompt 长度达到几千甚至上万 token 时标准注意力实现不仅慢还很容易把显存打爆。如果业务场景只关心局部上下文比如长文档问答、流式对话、代码补全滑动窗口注意力Sliding Window Attention是一个很经典的选择。而 FlashAttention 又是让滑动窗口注意力真正“跑起来、跑得快”的关键加速器。这篇文章围绕一个具体问题展开FlashAttention 如何加速滑动窗口注意力的 prefill我会从 prefill 的概念讲起逐步拆解 FlashAttention 的分块计算思路再结合朴素 PyTorch 实现与 flash-attn 库的实际调用方式说明窗口机制与 IO 优化为什么是天然匹配的。最后还会给出复杂度对比、常见坑点和工程建议。1. 背景与核心概念1.1 大模型推理的两个阶段prefill 与 decode接触过大模型推理部署的开发者大概率都听过这两个词prefill 和 decode。它们是大模型生成回答时天然存在的两个计算阶段。prefill预填充阶段用户输入 prompt 后模型一次性并行处理整个输入序列计算每个输入 token 的 hidden state同时生成第一轮 KV cache。这个阶段的特点是并行度高、计算密度大通常表现为长序列上的大矩阵乘。decode解码阶段模型逐 token 自回归生成输出。每步只处理一个新 token却要读取之前所有 token 的 KV cache因此这个阶段通常是访存密集memory-bound的。在长 prompt 场景下prefill 的耗时往往占整体首 token 时延TTFT的绝大部分。比如一个 4K token 的 prompt如果一次性全量计算注意力那么注意力分数矩阵是 4K × 4K显存占用和计算量都非常可观。这也是为什么很多推理框架vLLM、TensorRT-LLM 等会专门优化 prefill甚至把长 prompt 拆成多个 chunk 分段计算。1.2 标准注意力为什么慢N² 矩阵是瓶颈我们先看标准缩放点积注意力Scaled Dot-Product Attention的计算公式Attention(Q, K, V) softmax(QK^T / sqrt(d)) V假设序列长度为 Nhead 维度为 d。计算 Attention 时需要先算出 QK^T得到形状为(N, N)的注意力分数矩阵 S。这个 S 矩阵有两大问题显存占用是 O(N²)。N4096 时S 是 4096×4096 的矩阵N32768 时S 就有 10 亿个元素即使是 fp16 也要 2GB 显存。prefill 阶段 batch 内一般还有多条请求乘上 batch 和 head 数之后增长更快。HBM显存访问量巨大。标准实现里S 矩阵要写回显存、再读出来做 softmax、softmax 结果 P 还要再写回、再读出来乘 V。每一次写和读都是 N² 级别的显存带宽开销。GPU 算力很强但显存带宽是相对稀缺的N² 级别的搬运会直接把计算速度拖下来。所以标准注意力的核心瓶颈是中间结果 N×N 矩阵被反复物化materialize到显存中。这不仅占显存也浪费带宽。1.3 滑动窗口注意力限制每个 token 只看窗口内历史滑动窗口注意力的想法非常直观当前 query 位置 $i$ 不需要关注整个历史序列只需要关注左侧最近 $W$ 个 token。于是注意力分数的带状结构代替了原来的稠密三角结构。以 Mistral 7B 为代表的一批模型就使用了滑动窗口注意力。Mistral 的窗口大小设置为 4096配合 32 层 Transformer 层理论上信息可以通过层间传递让更远处的 token 间接影响当前 token所以有效感受野可以大于单个窗口。滑动窗口注意力的收益也很明显计算量从 O(N²·d) 降到 O(N·W·d)。显存占用从 O(N²) 降到 O(N·W)。在推理阶段KV cache 也可以只保留窗口内的历史进一步降低显存。但这里有一个容易被忽略的问题朴素实现里加 mask 只是把不想要的位置“挡住”并不会减少计算量。你依然先算了完整的 QK^T再把窗口外的位置填成负无穷。也就是说算法复杂度下来了但工程实现如果不做剪枝实际的 GPU 计算量一点没少。1.4 三个概念如何组合成今天的主题把上面三个概念串起来问题就清晰了prefill 阶段需要快速处理长 prompt滑动窗口注意力可以把算法复杂度从 N² 降到 N·W但要让这个复杂度的降低真正反映到 GPU 速度上不能用“先算完整矩阵再 mask”的笨办法而是要在内核层面直接跳过窗口外的块。FlashAttention 恰好提供了这种能力它把注意力计算分块化每个 query 块只加载需要的 K/V 块当配合滑动窗口时窗口外的 K/V 块可以直接跳过不需要物化 N×N 分数矩阵也不需要计算窗口外位置的分数。2. FlashAttention 加速的核心原理2.1 从 HBM 与 SRAM 的存储层级说起要理解 FlashAttention先要理解 GPU 的存储层级。现代 GPU比如 A100、H100内部有两种关键存储HBMHigh Bandwidth Memory显存容量大几十 GB但带宽相对有限。SRAMStatic RAM芯片内的高速缓存/共享内存容量很小通常几十 KB 到几百 KB但带宽极高、延迟极低。标准注意力的问题在于每一步中间结果都要经过 HBM把高带宽的 SRAM 浪费了。FlashAttention 的核心思路是IO-aware让计算尽量在 SRAM 中完成减少 HBM 读写次数。FlashAttention 是 Dao 等人 2022 年提出的精确注意力算法不是近似方法它把 Q、K、V 切成小块每次只把一个小块加载到 SRAM在 SRAM 内完成分数计算、softmax、加权求和再把结果写回 HBM。整个过程中N×N 的注意力分数矩阵从未被完整物化到显存。2.2 分块 Tiling让注意力矩阵不出 SRAM分块Tiling是 FlashAttention 最核心的工程技巧。假设一个块的大小为 B那么把 Q 分成 N/B 个 query 块把 K、V 分成 N/B 个 key/value 块对于每一个 query 块遍历所有 key/value 块在 SRAM 内计算这个子块对应的分数矩阵(B, B)更新 softmax 统计量并累加结果。之所以能这样做是因为注意力计算的核心操作可以拆成“矩阵乘 逐行 softmax 矩阵乘”的组合。只要每个子块都能独立更新输出累加器就能把大矩阵拆成小矩阵流式计算。这里有一个关键点FlashAttention 不是简单地“分块然后拼起来”因为 softmax 有一个全局的归一化项。如果先算某个块的 softmax再用另一个块的统计量去缩放结果是错的。所以还需要下一节的在线 Softmax 技巧。2.3 在线 Softmax 与重缩放在标准实现中softmax 需要知道整行的最大值然后对所有元素做指数归一化softmax(x)_i exp(x_i - max(x)) / sum(exp(x - max(x)))但 FlashAttention 是逐块处理的处理第一块时并不知道后续块里的最大值。怎么保证最终结果和全局 softmax 一致答案是在遍历 K/V 块的过程中维护两个统计量当前行最大值 mrunning max当前行指数和 lrunning sum相当于 softmax 分母。每处理一个新块先计算该块内的行最大值m_new。然后如果m_new比之前的 m 更大说明之前的累加器需要整体缩小缩放因子是exp(m - m_new)当前块的exp(s - m_new)可以直接参与累加分母 l 也按同样的方式缩放并累加。遍历完所有块后最终输出为acc / l。这个“在线 softmax”技巧最早出现在 softmax 数值稳定的工程实现中FlashAttention 把它和分块融合了起来既保证了数值精度又只对每个块做一次读取。2.4 FlashAttention 相比标准实现的 IO 差异从 HBM 访问量的角度看标准注意力需要对注意力矩阵做多次读写HBM 访问量大约是 O(N² N·d)其中 N² 项来自矩阵 S 和 P 的写回与重读FlashAttention 通过分块把一个 query 块需要访问的 K/V 块流式读完块内的中间结果不落显存。按块大小 B 来估算HBM 访问量近似为 O(N·d N²·d/B)。当 B 取到 64 或 128 时N² 带来的带宽压力会被显著摊薄。要注意FlashAttention 并不是减少计算量计算复杂度仍然是 O(N²·d)。它减少的是 HBM 搬运量和显存占用让 GPU 的矩阵乘单元比如 Tensor Core可以满负荷工作。这也是为什么在很多官方评测里FlashAttention 能在不损失精度的前提下做到数倍加速。3. 滑动窗口注意力的结构优势3.1 注意力矩阵的带状结构滑动窗口注意力下的注意力矩阵是什么样子为了直观我们画一个小例子。假设序列长度 N6窗口大小 W3并且使用因果掩码那么每个 query 位置 i 只能看到[i-2, i]这个范围内的 keyi - j 3 且 j i。用代码块表示合法的注意力位置X 表示可以看到X X . . . . X X X . . . . X X X . . . . X X X . . . . X X X . . . . X X可以看到合法位置集中在主对角线附近形成一条“带”。如果序列更长、块更大这条带在块级别看就是“块带状”的一个 query 块只需要访问窗口范围内的少数几个 K/V 块。3.2 朴素实现的问题mask 只挡计算不省算力很多初学者会写类似下面这样的代码先算出完整的 QK^T然后构造一个 mask把窗口外的位置设为-inf再 softmax。这种做法在功能上没问题在性能上却有两个浪费计算浪费窗口外的位置仍然参与了 QK^T 的矩阵乘。比如窗口 W512、序列长度 N4096实际只需要算 1/8 的分数但朴素实现把 8 倍的计算量全算了。显存浪费N×N 的分数矩阵还是会完整出现在显存中mask 之前它就已经占用了大量 HBM 带宽和显存。所以“滑动窗口 朴素 mask”只能说在理论上降低了复杂度在工程上并没有兑现。3.3 块级跳过与元素级掩码FlashAttention 的天然匹配FlashAttention 的分块方式与滑动窗口的带状结构非常契合。具体来说在遍历 K/V 块时可以通过简单的索引判断来决定如果某个 K/V 块整体位于当前 query 块的窗口左侧之外直接跳过如果某个 K/V 块整体位于当前 query 块的右侧因果之外直接结束循环如果 K/V 块与窗口部分重叠则对这个块内部做元素级掩码。块级跳过的判断只是整数比较几乎没有额外开销。而正因为窗口外的块被跳过了FlashAttention 滑动窗口的计算量才真正降到了 O(N·W·d)。这也是它相比“稀疏矩阵库”更优雅的地方滑动窗口的稀疏模式非常规则不需要存储复杂的稀疏索引直接在块循环里用索引算术就能表达。4. 完整代码实战4.1 环境说明本文的示例以 PyTorch 和 flash-attn 库为主。具体版本需要根据你的环境调整以下只是示例思路Python 3.10PyTorch 2.x建议用官方镜像或 pip 安装CUDA 11.8 或更高版本flash-attn 2.x 版本库flash-attn 的安装有时会比较费劲因为它是编译安装的。如果你的 CUDA 版本与 PyTorch 版本不一致很容易遇到导入报错。建议先确认torch.version.cuda与nvidia-smi中的 CUDA 版本互相兼容再参考官方安装说明。下面先用一个朴素 PyTorch 实现建立正确性基线。4.2 朴素 PyTorch 实现滑动窗口注意力我们先实现一个“正确但慢”的版本。它先算完整分数矩阵再叠加因果掩码和窗口掩码。import torch import torch.nn.functional as F def sliding_window_attention_naive(q, k, v, window_size): 朴素滑动窗口注意力正确性参考实现 q, k, v: shape (batch, heads, seq_len, head_dim) window_size: 每个 query 最多可以看到左侧多少个 token b, h, n, d q.shape scale d ** -0.5 # 完整分数矩阵 (b, h, n, n) scores torch.matmul(q, k.transpose(-2, -1)) * scale # 构造掩码因果 窗口 rows torch.arange(n, deviceq.device).view(-1, 1) cols torch.arange(n, deviceq.device).view(1, -1) causal rows cols # query i 只能看 key j i window (rows - cols) window_size # i - j window_size mask causal window scores scores.masked_fill(~mask, float(-inf)) attn torch.softmax(scores, dim-1) out torch.matmul(attn, v) return out这里masked_fill把窗口外和未来位置都置为-infsoftmax 后这些位置的概率为 0。这个实现适合验证逻辑不适合做性能测试。4.3 使用 flash-attn 库开启 window_sizeflash-attn 库提供的flash_attn_func支持window_size参数。注意它的输入布局与 PyTorch 原生注意力不同一般是(batch, seq_len, num_heads, head_dim)。from flash_attn import flash_attn_func # 输入布局(batch, seq_len, num_heads, head_dim) # 示例batch2, seq_len1024, heads8, head_dim64 q torch.randn(2, 1024, 8, 64, dtypetorch.float16, devicecuda) k torch.randn(2, 1024, 8, 64, dtypetorch.float16, devicecuda) v torch.randn(2, 1024, 8, 64, dtypetorch.float16, devicecuda) # 滑动窗口 因果左侧窗口 256右侧为 0 W 256 out flash_attn_func( q, k, v, dropout_p0.0, softmax_scaleNone, causalTrue, window_size(W, 0), )这里有几个参数要解释causalTrue开启因果掩码query 只能看到当前位置及之前的位置。window_size(W, 0)表示左侧窗口大小为 W右侧窗口大小为 0。在因果场景下右侧本来就看不到任何 token所以右侧设为 0 是合理且常见的写法。window_size(-1, -1)表示不限制窗口等价于完整注意力。需要特别提醒window_size的语义和实现版本有关不同版本的 flash-attn 对边界情况的处理可能有细微差异。在使用前最好先查看你安装版本的函数签名和文档。4.4 内核主循环伪代码块循环与窗口判断为了理解 FlashAttention 到底在 GPU 里做了什么我们来看一个简化的伪代码。它不是可以直接运行的 Python 代码但能清晰表达块级跳过的逻辑。# 伪代码FlashAttention 滑动窗口前向内核逻辑仅供理解 # Q/K/V/O 均为 (N, d)按块大小 B 切分 # W 为左侧窗口大小 def flash_attn_sliding_window_kernel(Q, K, V, O, B, W): for q_block_id in range(N // B): q load_q_block(q_block_id) # (B, d) q_start q_block_id * B m fill(-inf, B) # 行最大值 l zeros(B) # softmax 分母 acc zeros(B, d) # 输出累加器 for kv_block_id in range(N // B): k_start kv_block_id * B # 块级剪枝 1整个 K/V 块都在窗口左侧之外 if k_start B q_start - W: continue # 块级剪枝 2整个 K/V 块都在当前 query 块右侧因果 if k_start q_start B: break k load_k_block(kv_block_id) # (B, d) v load_v_block(kv_block_id) # (B, d) # 子块分数矩阵只存在于 SRAM s q k.T * scale # (B, B) # 元素级掩码块内处理窗口边界 s mask_block(s, q_start, k_start, W) m_new maximum(m, row_max(s)) p exp(s - m_new) alpha exp(m - m_new) l alpha * l row_sum(p) acc alpha * acc p v m m_new o acc / l write_o_block(q_block_id, o)两个if判断就是滑动窗口与 FlashAttention 结合的精华第一个if跳过窗口左侧之外的块第二个if利用因果性一旦 K/V 块完全在未来就直接跳出内层循环。这样每个 query 块实际访问的 K/V 块数量从 N/B 降到了大约 W/B计算量和 HBM 读取量都随窗口大小线性缩减。4.5 正确性验证思路把朴素实现作为 baseline与 flash-attn 的输出做对比可以确认自己的调用方式是否正确。# 随机输入对比朴素实现与 flash_attn 输出 torch.manual_seed(0) b, n, h, d 2, 1024, 8, 64 q torch.randn(b, n, h, d, dtypetorch.float16, devicecuda) k torch.randn(b, n, h, d, dtypetorch.float16, devicecuda) v torch.randn(b, n, h, d, dtypetorch.float16, devicecuda) W 256 out_fa flash_attn_func(q, k, v, dropout_p0.0, causalTrue, window_size(W, 0)) # 朴素实现需要在 head 维度上调整布局 out_naive sliding_window_attention_naive( q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), W ) err (out_naive.transpose(1, 2) - out_fa).abs().max().item() print(max abs error:, err)正常情况下fp16 下的最大绝对误差会在 1e-2 量级甚至更小。如果发现误差很大优先检查两个方向window_size的左右语义是否用反了是否把causalTrue错误地设置成了False导致未来位置的掩码不一致。5. 复杂度对比与性能分析5.1 计算量与 HBM 访问量的公式对比为了说清楚收益我们用 N 表示序列长度W 表示窗口大小d 表示 head 维度B 表示分块大小。这里只列出主要项忽略常数系数。标准注意力计算量O(N²·d)HBM 访问量O(N² N·d)主要来自 N×N 分数矩阵的写回与重读中间显存O(N²)朴素滑动窗口 mask计算量仍然是 O(N²·d)因为 mask 之前已经算了完整 QK^THBM 访问量O(N² N·d)中间显存O(N²)FlashAttention 滑动窗口计算量O(N·W·d)因为窗口外的块被跳过HBM 访问量O(N·d N·W·d/B)块内中间结果不落显存中间显存O(N) 级别只有输入输出需要驻留可以看到真正把“算法复杂度下降”转化为“GPU 计算量下降”的关键是 FlashAttention 的块级跳过。如果窗口 W 远小于 N那么 N·W 相比 N² 的收益会非常可观。5.2 对比表格实现方案计算量主要 HBM 开销中间矩阵是否物化标准全注意力O(N²·d)O(N²)是N×N标准 滑窗 maskO(N²·d)O(N²)是N×N稀疏注意力理论实现O(N·W·d)O(N·W)可能是块稀疏矩阵FlashAttention 全注意力O(N²·d)O(N²·d/B)否只保留 O(N)FlashAttention 滑窗O(N·W·d)O(N·W·d/B)否只保留 O(N)这里 B 受 SRAM 容量限制实践中通常取 64 或 128。所以“FlashAttention 滑窗”不仅把计算量从 N² 级降到了 N·W 级还把 HBM 访问量进一步除以 B双管齐下。5.3 为什么 prefill 阶段收益最大前面提到prefill 阶段处理的是整个 prompt所有输入 token 并行计算。在这个阶段序列长度 N 很大N² 矩阵的代价被放大计算是“多行 query × 全部 KV”非常适合分块并行FlashAttention 的分块逻辑可以把 N×N 的计算切成小矩阵乘充分利用 GPU 的 Tensor Core。因此在长 prompt 的 prefill 上开启滑动窗口 FlashAttention 通常能看到两个明显效果显存占用显著下降prefill 耗时显著缩短。而短 prompt 场景比如 N 只有几十到一两百N² 本来就不大此时收益不明显甚至可能因为内核启动和分块开销显得没多少提升。5.4 decode 阶段的不同点decode 阶段每次只有一个新 token相当于一个 query 行需要读取全部或窗口内KV。这个阶段本质上是从 HBM 读取 KV cache 的访存操作计算量很小。FlashAttention 对 decode 也有优化但瓶颈不再是 N² 矩阵而是 KV cache 的读取带宽。滑动窗口在这里的好处是KV cache 可以只保留窗口内的部分显存占用从 O(N·d) 降到 O(W·d)。所以一个常见结论是prefill 阶段靠 FlashAttention 减少冗余计算与 HBM 搬运decode 阶段靠窗口裁剪 KV cache 减少访存量。两者结合才能让滑动窗口模型在长序列推理上达到理想性能。6. 常见问题与排查思路问题现象常见原因解决思路导入 flash_attn 报 ImportError编译安装时的 CUDA 版本与运行时不一致确认torch.version.cuda与驱动兼容优先使用官方预编译 wheel 或源码编译设置了 window_size 但结果与预期不符左右窗口语义弄反或未开启 causal用随机输入对比朴素实现先测window_size(-1,-1)再逐步收窄窗口输出与朴素实现误差很大mask 逻辑不一致或布局错误检查(batch, seq, head, dim)与(batch, head, seq, dim)是否混用长 prompt prefill 仍然 OOM没有真正走到 flash 内核或 head 维度过大打印内核日志确认使用 flash考虑结合 chunked prefill 分段处理短 prompt 加速不明显N 太小N² 开销本身可忽略这是正常现象短 prompt 直接使用完整注意力即可训练时反向传播很慢FlashAttention backward 需要重计算分数矩阵这是预期行为可通过调整块大小或优化 GPU 型号缓解除了表格里的问题还有一个比较隐蔽的坑滑动窗口必须和模型训练时的注意力模式一致。如果模型是用完整注意力预训练的推理时突然切成滑动窗口远距离依赖会丢失输出质量会明显