ARTICLE DETAIL

资讯详情

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

全注意力机制为何成为大模型推理的性能瓶颈?

全注意力机制为何成为大模型推理的性能瓶颈? 不少同学在跑大模型推理时应该都有过这种体感模型生成第一个字很快越往后越慢上下文特别长的时候每个新 token 都要等很久。网上查资料得到的答案往往是“因为注意力机制是 O(n²) 的复杂度”但很少有人能讲清楚这个 O(n²) 究竟消耗在哪里为什么是“每个新词都要重翻百万页记录”。这篇文章是 Kimi Linear 核心原理系列的第一篇我们先把问题的最底层讲透全注意力机制到底贵在哪儿。我会从注意力机制的基本原理讲起配合可运行的代码示例一步步拆解全注意力在训练和推理阶段的计算开销并解释 KV Cache 为什么能缓解问题、为什么长上下文下依然不够最后再引出线性注意力这类高效方案的优化思路。适合刚接触大模型原理的开发者也适合已经会用模型但想深入理解推理性能瓶颈的工程师。1. 背景先理解全注意力是什么注意力机制Attention最早在机器翻译任务中大规模应用后来随着 Transformer 架构的提出成为大模型的核心组件。简单来说注意力机制让模型在处理当前词时能够有选择地“关注”输入序列中的其他词而不是像传统循环神经网络那样只依赖上一步的隐状态。全注意力Full Attention通常指标准 Transformer 中使用的缩放点积注意力Scaled Dot-Product Attention。它的特点在于序列中的每一个 token都要与序列中的所有token 计算关联度。这里“所有”两个字就是计算开销的根源。1.1 从“翻聊天记录”理解注意力我们可以用一个很生活化的比喻来理解注意力机制。假设你是一个客服面前摆着厚厚的聊天记录。用户问了你一个新问题你需要回答。你可能需要翻看之前的对话找到相关信息。全注意力机制做的事情是每回答一个词就把前面所有聊天记录从头到尾重新翻一遍。第一次回答时记录只有一页翻起来很快。第 100 个词时记录累积到 100 页为了生成这 1 个词你至少要把这 100 页扫描一遍。第 1000 个词时你面前已经有 1000 页记录而你还是需要把每一页都过目一遍才能生成下一个词。这就是标题里“每个新词都要重翻百万页记录”的含义。百万页记录越堆越多翻一遍的成本越来越高生成新词的速度自然就越来越慢。1.2 注意力机制的正式定义从数学上看缩放点积注意力可以写成Attention(Q, K, V) softmax(QK^T / sqrt(d_k)) V其中QQuery查询向量代表“当前正在处理的词”可以理解为“我在找什么”。KKey键向量代表“序列中每个词的标签”可以理解为“我这里有什么”。VValue值向量代表“序列中每个词的实际内容”可以理解为“我能提供什么信息”。d_k 是 Q 和 K 的维度除以 sqrt(d_k) 是为了防止点积结果过大导致 softmax 梯度消失。先计算 Q 和所有 K 的点积得到当前词与每个历史词的相关性分数再经过 softmax 转换成概率分布最后把这些概率作为权重对所有 V 做加权求和得到当前词的注意力输出。整个计算过程中QK^T 需要矩阵乘法结果矩阵的大小是“序列长度 × 序列长度”。这个“序列长度 × 序列长度”的矩阵就是全注意力复杂度的直接来源。2. 全注意力贵在三个维度很多人提到注意力复杂度只记住了一个“O(n²)”。但在实际工程中我们发现全注意力贵在三个层面缺一不可。2.1 计算量n² 的矩阵乘法给定一个长度为 n 的输入序列QK^T 的计算结果是一个 n × n 的矩阵。矩阵乘法: [n × d] × [d × n] [n × n]也就是说每生成一个 token都需要做一次规模与历史长度平方成正比的矩阵运算。如果现在的输入长度是 1000那么注意力矩阵就是 1000 × 1000如果长度变成 10000矩阵就是 10000 × 10000计算量增长了 100 倍。这就是为什么长上下文场景下生成速度下降非常明显。2.2 内存占用每层都要存 n² 的注意力矩阵Transformer 通常有几十层每一层都有自己的注意力头。在训练阶段前向传播需要保存每一层的注意力矩阵用于反向传播时计算梯度。一个长度为 n 的序列一层一头的注意力矩阵大小就是 n²。举个例子序列长度 4096注意力矩阵是 4096 × 4096一个 float32 数值占 4 字节那么单个注意力矩阵就占 67 MB。如果模型有 32 层、每层多个注意力头内存开销会迅速膨胀。2.3 内存带宽推理阶段真正的瓶颈很多人以为推理阶段的瓶颈是计算量但实际上现代 GPU 上更常见的瓶颈是内存带宽Memory Bandwidth。原因在于虽然我们不需要在推理时保存 n² 的注意力矩阵用于反向传播但每生成一个 token仍然需要把完整的 K 和 V 矩阵从显存中读出来参与计算。我们来算一笔账假设上文长度是 10000。每一层的 K 和 V 矩阵大小为 10000 × d其中 d 是注意力头维度。假设 d 128那么每层每头的 K 和 V 合起来是 10000 × 128 × 2 2.56M 个数值。一个 float16 数值占 2 字节那么每层每头需要读取约 5 MB 数据。如果是 32 层、每层 32 个头单次生成需要读取的数据量是 5 MB × 32 × 32 5.12 GB。也就是说模型每生成一个 tokenGPU 就要搬运约 5 GB 的显存数据。这个数据搬运的时间往往会超过实际计算的时间。这也是为什么在实际推理中我们经常看到 GPU 利用率不高但 token 生成速度就是上不去。3. 用一段代码看清全注意力的“重读”代价这部分我们通过一个最小实验直观感受全注意力在生成阶段的开销变化。3.1 最小注意力实现先用 NumPy 实现一个标准的缩放点积注意力import numpy as np def scaled_dot_product_attention(q, K, V): 标准缩放点积注意力 参数: q: 当前 token 的查询向量, 形状为 (d_k,) K: 历史所有 token 的键矩阵, 形状为 (n, d_k) V: 历史所有 token 的值矩阵, 形状为 (n, d_v) 返回: output: 当前 token 的注意力输出, 形状为 (d_v,) weights: 当前 token 对所有历史 token 的注意力权重, 形状为 (n,) # 计算当前 query 与所有 key 的点积 scores q K.T # 形状: (n,) # 缩放 d_k q.shape[0] scores scores / np.sqrt(d_k) # softmax 归一化 weights np.exp(scores - np.max(scores)) weights weights / weights.sum() # 形状: (n,) # 加权求和 output weights V # 形状: (d_v,) return output, weights这个函数模拟的就是“每生成一个词都要翻一遍所有历史 token”的过程。q 是当前查询向量K 和 V 是所有历史 token 的键和值。注意这里 K 和 V 的规模 n 越大scores 的计算量就越大。3.2 模拟生成过程中的注意力开销接下来模拟一个自回归生成过程假设输入 prompt 已经有prompt_len个 token我们还要继续生成new_tokens个新 token。def simulate_full_attention_generate(prompt_len, new_tokens, d64): 模拟自回归生成场景下全注意力机制的累计计算量 计算量统计方式QK^T 矩阵乘法需要 n*d 次乘法近似 这里为了直观比较数量级用浮点运算次数粗略表示 返回: step_flops: 每一步生成所需的 FLOPs 列表 step_flops [] # 每生成一个新 token历史长度增加 1 for step in range(new_tokens): n prompt_len step # 当前需要 attend 的历史 token 数 # QK^T 计算量约 2 * n * d忽略 softmax 和加权求和部分 flops_qk 2 * n * d # scores 与 V 加权求和计算量约 2 * n * d flops_weighted_sum 2 * n * d # 还有 softmax 内部的对历史 n 个分数的归一化 flops_softmax 3 * n step_flops.append(flops_qk flops_weighted_sum flops_softmax) return step_flops运行这个模拟prompt_len 100 new_tokens 2000 step_flops simulate_full_attention_generate(prompt_len, new_tokens) # 观察第 1 步、第 100 步、第 1000 步和第 2000 步的开销 for idx in [0, 99, 999, 1999]: print(f生成第 {idx 1} 个新词历史长度 {prompt_len idx}) print(f 该步注意力计算量约: {step_flops[idx]:.2e} FLOPs)预期输出如下生成第 1 个新词历史长度 100 该步注意力计算量约: 2.57e04 FLOPs 生成第 100 个新词历史长度 199 该步注意力计算量约: 5.14e04 FLOPs 生成第 1000 个新词历史长度 1099 该步注意力计算量约: 2.83e05 FLOPs 生成第 2000 个新词历史长度 2099 该步注意力计算量约: 5.39e05 FLOPs可以看到生成第 2000 个词时单步的计算量大约是第 1 个词时的 20 倍。而且随着生成继续这个差距会越来越大。3.3 为什么感觉模型越跑越慢上面的模拟还只是朴素的全注意力计算。如果把所有层的计算量都加起来差异会更明显。我们稍微扩展一下代码统计生成到不同阶段时的累计开销import matplotlib.pyplot as plt def compute_cumulative_flops(prompt_len, new_tokens, num_layers32, num_heads32): 计算多层多头的累计注意力计算量 step_flops simulate_full_attention_generate(prompt_len, new_tokens) # 每层、每个注意力头都有一次注意力计算 single_step_total [f * num_layers * num_heads for f in step_flops] # 累计开销 cumulative [] running_sum 0 for f in single_step_total: running_sum f cumulative.append(running_sum) return cumulative # 计算累计浮点运算量 cumulative_flops compute_cumulative_flops(100, 2000) # 输出几个关键节点 for idx in [9, 99, 499, 999, 1999]: print(f生成到第 {idx 1} 个词累计注意力计算量约: {cumulative_flops[idx]:.2e} FLOPs)预期输出生成到第 10 个词累计注意力计算量约: 1.14e08 FLOPs 生成到第 100 个词累计注意力计算量约: 1.18e10 FLOPs 生成到第 500 个词累计注意力计算量约: 2.98e11 FLOPs 生成到第 1000 个词累计注意力计算量约: 1.29e12 FLOPs 生成到第 2000 个词累计注意力计算量约: 5.29e12 FLOPs从这个模拟可以清楚看到生成的前 100 个词累计计算量只有 1.18e10而生成到第 1000 个词时已经累计到了 1.29e12到第 2000 个词时累计到了 5.29e12。越往后每生成一个词需要额外付出的代价越高。这也解释了为什么大模型在长对话场景下响应速度会明显变慢。4. 工程上的优化KV Cache 是什么又为什么不够既然全注意力在生成时每次都要重新计算 K 和 V一个很自然的想法就是反正之前 token 的 K 和 V 不会变为什么不把它们缓存下来只缓存一次以后直接复用这就是 KV Cache 的由来。4.1 KV Cache 的基本原理在自回归生成中每一步只会新增一个 token。已经生成的 token 不再变化因此它们在所有层的 K 和 V 已经固定。如果不缓存每一步生成时都要重新计算所有历史 token 的 K 和 V如果缓存只需要计算当前新 token 的 K 和 V然后追加到缓存里。下面用代码对比一下缓存和不缓存的计算差异def no_cache_kv_compute(prompt_len, new_tokens, layers32): 不缓存 K/V每一层每一步都重新计算全部历史 token 的 K 和 V 粗略用 token 数表示计算规模 total 0 for step in range(new_tokens): current_len prompt_len step total current_len * 2 * layers return total def with_cache_kv_compute(prompt_len, new_tokens, layers32): 使用 KV Cache只计算 prompt 阶段和历史生成阶段的 K/V 新 token 的 K/V 在每步新增时计算一次后续复用 # prompt 阶段需要计算全部 prompt_len 个 token 的 K/V total prompt_len * 2 * layers # 生成阶段每步新增 1 个 token 的 K/V for step in range(new_tokens): total 2 * layers # 当前一步新增的 1 个 token return total prompt_len 1000 new_tokens 2000 layers 32 no_cache_total no_cache_kv_compute(prompt_len, new_tokens, layers) with_cache_total with_cache_kv_compute(prompt_len, new_tokens, layers) print(f不缓存 K/V累计 K/V 计算规模: {no_cache_total:.2e}) print(f使用 KV Cache累计 K/V 计算规模: {with_cache_total:.2e}) print(f优化比例: {no_cache_total / with_cache_total:.2f} 倍)预期输出不缓存 K/V累计 K/V 计算规模: 1.92e08 使用 KV Cache累计 K/V 计算规模: 2.56e05 优化比例: 750.00 倍也就是说KV Cache 能把 K 和 V 的重复计算量减少几百倍。这也是为什么现代推理框架普遍使用 KV Cache。4.2 为什么有了 KV Cache 还不够既然 KV Cache 这么好为什么长上下文下还会慢原因有两点第一KV Cache 只是省去了 K 和 V 的重复计算但注意力计算本身仍然需要遍历全部历史 token。Q 仍然要和所有历史的 K 做点积然后对所有历史的 V 做加权求和。这段计算量并不会因为缓存而减少。第二KV Cache 占用显存并且随着文本长度线性增长。我们前面算过如果模型很大、层数很多单次生成就需要从显存读取很大的 KV Cache。当上下文长度达到几十万甚至上百万 token 时KV Cache 的显存占用会非常夸张。算一笔账def estimate_kv_cache_size(num_layers, num_heads, head_dim, seq_len, dtype_bytes2): 估计 KV Cache 大小 # 每个 token 每一层需要保存 K 和 V各 num_heads * head_dim 个数值 kv_per_token_per_layer 2 * num_heads * head_dim # 总字节数 total_bytes kv_per_token_per_layer * num_layers * seq_len * dtype_bytes return total_bytes # 以 7B 规模模型常见配置为例 num_layers 32 num_heads 32 head_dim 128 seq_len 100000 # 10万上下文 cache_size estimate_kv_cache_size(num_layers, num_heads, head_dim, seq_len) print(f10万上下文时 KV Cache 约: {cache_size / 1024**3:.2f} GB) seq_len 1000000 # 100万上下文 cache_size estimate_kv_cache_size(num_layers, num_heads, head_dim, seq_len) print(f100万上下文时 KV Cache 约: {cache_size / 1024**3:.2f} GB)预期输出10万上下文时 KV Cache 约: 31.25 GB 100万上下文时 KV Cache 约: 312.50 GB这只是单个请求的 KV Cache而且只算了 fp16 精度。如果是并发多用户场景显存压力会成倍上升。所以 KV Cache 虽然大幅减少了重复计算但面对百万级上下文仍然不是一个终极方案。5. 从全注意力到 Linear核心思路是去掉 n²既然全注意力贵的根源是“每个 token 都要和所有 token 计算关联”那么一个自然的优化方向就是能不能不要遍历所有历史 token这就是线性注意力Linear Attention等高效注意力机制要解决的问题。5.1 线性注意力的基本思想标准注意力的计算顺序是先计算 QK^T 得到 n×n 的相似度矩阵 再对 V 做加权求和线性注意力尝试改变这个顺序利用矩阵乘法的结合律把计算顺序调整成先计算 K^T V再和 Q 相乘。这里的关键在于K^T V 的结果是一个 d×d 的矩阵和序列长度 n 无关。这样复杂度就从 O(n²) 降到了 O(n)也就是线性复杂度。可以这样理解标准注意力是“每一页都翻一遍找出和当前问题相关的页码”。线性注意力是“先看目录和摘要提炼出关键信息再基于当前问题快速查一次”。这样处理大幅减少了对历史记录的重复翻阅。5.2 这类方案在长上下文中的价值线性注意力的核心优势体现在两个地方单步生成开销稳定。无论历史上下文增长到多长每一步生成的计算量只和固定维度的 d 有关而不是和当前序列长度 n 有关。显存占用更可控。KV Cache 大小不再随序列长度线性增长某些变体能做到固定大小的状态表示这对于超长上下文部署非常有吸引力。不过也要注意线性注意力并不是免费的午餐。把 softmax 注意力改写成线性形式后往往会牺牲一部分表达能力。很多论文和工程实现会在计算效率与效果之间做平衡例如保留局部窗口注意力、加入滑动窗口机制或者采用门控机制增强状态更新能力。具体哪一种方案更优取决于实际任务。5.3 Kimi Linear 这个方向在做什么从命名上看Kimi Linear 属于高效注意力机制的探索方向这个系列会围绕线性注意力原理、实现、训练与推理优化展开。作为系列第一篇这一篇先把全注意力的代价彻底讲清楚。需要特别说明的是不同模型和框架对线性注意力的落地方式并不完全一致。有的方案通过内核融合优化计算过程有的方案通过状态空间模型SSM近似长程依赖有的方案用稀疏注意力替代全量注意力。我们不能想当然地认为所有 Linear 方案都使用同一种数学变换。在实际阅读相关代码和论文时建议以具体的开源实现和官方文档为准重点关注它修改的是注意力计算哪个环节以及复杂度到底降到了什么量级。6. 如何观察自己模型里的注意力开销理解理论之后我们需要在工程上能够观察和量化注意力开销。这里给出几个常用的观察思路。6.1 从推理日志看首 token 时延与 Decode 时延大多数推理框架都会输出两类指标Prefill首 token 时延处理输入 prompt 并生成第一个 token 的时间。Decode逐 token 时延生成后续每个 token 的平均时间。在标准 Transformer 中Prefill 阶段因为输入很长计算量主要花在 QK^T 大矩阵乘上Decode 阶段虽然每一步只算一个 token 的注意力但需要读取全部历史的 KV Cache因此内存带宽压力更大。如果你发现随着对话轮次增加Decode 时延显著上升大概率就是全注意力的“重读历史”成本在上升。6.2 用 PyTorch 的 profiling 工具观察以 PyTorch 为例可以用torch.profiler查看注意力计算占用的时间import torch import torch.nn.functional as F def attention_scores(q, k, v, maskNone): d_k q.size(-1) scores torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d_k) if mask is not None: scores scores.masked_fill(mask 0, float(-inf)) weights F.softmax(scores, dim-1) output torch.matmul(weights, v) return output with torch.profiler.profile(activities[torch.profiler.ProfilerActivity.CUDA]) as prof: with torch.no_grad(): for _ in range(10): q torch.randn(1, 32, 4096, 64, devicecuda) k torch.randn(1, 32, 4096, 64, devicecuda) v torch.randn(1, 32, 4096, 64, devicecuda) output attention_scores(q, k, v) print(prof.key_averages().table(sort_bycuda_time_total, row_limit10))通过 profiling 结果你可以直观看到矩阵乘法和 softmax 各自占用的时间从而定位到底是计算瓶颈还是访存瓶颈。6.3 粗略估算模型的理论 KV Cache 大小前面已经给出了估算公式实际项目中可以用它评估当前显存是否够用def kv_cache_gb(num_layers, num_heads, head_dim, seq_len, batch_size1, dtype_bytes2): kv_bytes_per_token 2 * num_heads * head_dim * num_layers * dtype_bytes total_bytes kv_bytes_per_token * seq_len * batch_size return total_bytes / 1024**3 # 示例8B 级模型配置 print(8B 模型 128K 上下文 KV Cache:, kv_cache_gb(32, 32, 128, 128 * 1024))这段代码适合在面试里展示也适合在部署前做显存规划。7. 常见问题与排查思路在实际学习和使用中经常遇到下面这些问题问题现象常见原因解决思路生成速度随对话长度明显下降全注意力每次都要遍历全部历史 tokenKV Cache 读取量线性增长评估是否需要超长上下文必要时改用高效注意力方案或精简历史GPU 利用率不高但生成很慢瓶颈在内存带宽反复读取大体积 KV Cache开启内核融合、使用 PagedAttention 等页面式 KV Cache 管理降低访存压力长上下文时显存不足KV Cache 占用过大估算 KV Cache 大小减少并发数考虑量化 K/V或使用阶段性清理策略Prefill 阶段很慢Prompt 很长QK^T 矩阵非常大对 prompt 做截断或摘要使用分块注意力和 Flash Attention 优化训练时显存溢出需要保存 n×n 注意力矩阵用于反向传播使用梯度检查点或改用非标准的稀疏/线性注意力结构换了高效注意力后效果下降线性化近似损失了部分全局依赖建模能力在短上下文任务上做对比实验评估是否对特定任务影响明显排查这类问题时推荐按下面顺序进行确认当前上下文长度和模型配置。用 profiling 工具确认是计算瓶颈还是访存瓶颈。计算当前 KV Cache 理论大小判断显存压力。对比不同上下文长度下的 Decode 时延。根据业务需要选择是否引入高效注意力机制。8. 最佳实践与工程建议结合项目经验这里给出几条工程建议。8.1 不要盲目追求超长上下文全注意力模型在短文本上表现很好但强行塞入超长文本代价是推理速度下降和部署成本上升。实际业务中可以考虑先对输入做分块、摘要或知识库检索把真正相关的片段取出来而不是把全部历史都丢给模型。推荐做法1. 对历史对话做滚动窗口裁剪保留最近 N 轮关键内容。 2. 对长文档做切块并建立索引只取和当前问题最相关的若干块。 3. 用摘要模型压缩前置上下文保留结构化关键信息。8.2 使用成熟的推理加速方案如果你使用的是标准 Transformer 架构并部署到生产环境建议优先使用已经验证过的推理加速方案Flash Attention减少 HBM 读写让注意力计算更高效。PagedAttention把 KV Cache 按页管理减少显存碎片和浪费。KV Cache 量化把 K/V 压缩为更低精度降低访存带宽压力。这些方案不改变模型结构能在现有模型上直接获得提速。8.3 在训练阶段就考虑长上下文的影响如果核心业务确实需要超长上下文建议在模型训练阶段就引入高效注意力机制而不是事后做推理优化。因为训练阶段如果仍使用全注意力超长序列会直接带来巨大的显存开销。同时在模型评估时要单独记录不同上下文长度下的推理性能指标防止模型效果不错但部署不可用。8.4 安全与权限边界最后补充一条通用工程提醒对模型进行测试和部署时确保数据来源合法获得相应授权。涉及用户对话数据和私有文档时要遵守数据最小化原则不要在本地或者云端长时间保存不必要的原始文本。模型推理服务暴露到公网前做好鉴权、限流和日志脱敏避免产生数据泄露风险。9. 下一步学习路线这一篇的核心目标是把“全注意力为什么贵”讲透。读完这篇文章你应该已经掌握了以下内容全注意力机制的核心公式和计算流程。全注意力在训练和推理阶段分别贵在哪里。KV Cache 的作用以及为什么有了缓存长上下文依然很慢。高效注意力如 Linear Attention为什么能降低复杂度。如果你发现自己对注意力机制还不太熟悉建议先补一补 Transformer 的基本结构理解 Q、K、V 三个向量各自的意义再看 Flash Attention 的论文和代码理解访存优化与计算优化的区别。后续的系列文章可以继续深入探讨线性注意力的数学推导、状态空间模型与注意力机制的关系、高效注意力结构在大规模训练中的稳定性问题、推理框架中 KV Cache 的内存管理细节等。建议你按顺序学遇到不明白的数学推导拿笔在纸上推一遍然后一定动手复现一个小实验效果会比只看文章好很多。技术学习就是这样理解了底层代价才能真正理解为什么新的优化方案值得关注。建议收藏这篇文章遇到长上下文性能问题时再回来看看相关的估算方法和排查思路。
返回列表