ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

全注意力为何昂贵?从O(n²)复杂度到KV Cache显存瓶颈

全注意力为何昂贵?从O(n²)复杂度到KV Cache显存瓶颈 写大模型长文本应用时很多人会产生一种直观的疑惑为什么模型生成一个词好像要把之前所有的上下文都重新“读”一遍如果上下文有几十万字这个代价就会变得非常夸张。最近“Kimi Linear”这个话题热度很高很多人把它看作线性注意力方向的一次讨论聚焦。但在此之前我们得先把一个问题彻底讲清楚——全注意力到底为什么贵。本文就从全注意力的计算原理、复杂度、显存占用与生成过程入手把“每个新词都要重翻百万页记录”这件事背后的数学和工程原因拆开讲清楚也为后续理解线性注意力方案打好基础。1. 全注意力机制到底在做什么1.1 用一句话理解注意力注意力机制可以让模型在生成某个词时自主决定“应该重点看输入中的哪些位置”。这个“看”并不是模糊的检索而是用数值打分的方式把当前词与上下文里每个词之间的关联程度计算出来。举个例子句子是小明早上没有吃早餐所以到了中午他觉得很饿当模型处理“饿”这个词时它需要把很高的注意力权重放在“没有吃早餐”上而不是每一个词都一视同仁。注意力机制做的事情就是为当前词与其他所有词建立一个权重分布然后按权重把上下文信息聚合起来喂给模型。这就是“全注意力”中的“全”字当前词要和序列中每一个历史位置都做计算没有任何省略。1.2 Q、K、V 三件套的本质Transformer 中的注意力公式为Attention(Q, K, V) softmax(Q * K^T / sqrt(d_k)) * V其中QQuery表示“当前查询”类似你发出的搜索请求。KKey表示“被匹配的键”类似资料库中的索引。VValue表示“实际内容”类似索引对应的正文内容。计算过程可以拆成三步当前 token 的 Q 与所有位置的 K 做点积得到所有位置的匹配分数。把这些分数除以 sqrt(d_k) 做缩放再走一个 softmax变成概率分布。用这个概率分布去加权求和所有位置的 V得到最终输出。这里的关键是第 1 步Q 必须与每一个 K 做点积。假设上下文里有 100 万个 token那么计算一次注意力就需要做 100 万次点积。这就像你要写一个新的词就必须先把之前 100 万页记录翻一遍重新判断每一页与当前词的相关性。1.3 全注意力不是“一次性读完全部”还有一个容易误解的地方有人觉得既然预训练阶段模型已经见过海量数据推理时是不是就不用重新算了并不是。推理生成阶段模型会在内存/显存中保留当前对话或输入的全部上下文。每生成一个新 token这个新 token 必须重新去计算它和所有历史 token 的注意力。即使 K 和 V 被缓存下来了仍然需要做新 Q 对所有历史 K 的点积计算以及所有历史 V 的加权求和。所以“百万页记录”不是指预训练语料而是指当前输入上下文的长度。上下文越长这个重新翻阅动作就越贵。2. 每个新词都要重翻百万页记录的数学根源2.1 自回归生成中的两个阶段大模型生成文本时通常采用自回归方式一个 token 接一个 token 地输出。整个过程分为两个阶段预填充阶段对输入的 prompt 并行计算一遍生成初始的 K 和 V。解码阶段每生成一个新 token执行一次前向计算并更新 KV Cache。全注意力的昂贵主要体现在解码阶段。假设当前序列长度是 n新 token 的 Q 需要与 n 个历史 K 做点积也就是一次要处理 n 个位置。序列长度 n 等于多少就要重翻多少页记录。2.2 时间复杂度是 O(n²)先看单个 token 的注意力计算点积部分Q 是 1×dK 是 n×d相乘需要 n×d 次乘法。softmax 部分需要处理 n 个分数复杂度 O(n)。加权求和部分n×d 次乘法。所以单个 token 的注意力计算复杂度大约是 O(n·d)。生成 n 个 token 后总计算复杂度约是 O(n²·d)。当 n 很大时n² 就是最核心的瓶颈。如果说预训练阶段同一时刻每个位置并行计算复杂度同样是 O(n²·d)但硬件利用率较高观感上不如推理阶段那么明显。到了长文本生成场景每步延迟都要真实暴露出来O(n²) 的代价就非常压手。下面给出一个复杂度对比表计算阶段计算规模复杂度是否随序列增长加速恶化单 token 注意力1×n 点积O(n·d)线性增长整个序列注意力n×n 矩阵O(n²·d)平方增长KV Cache 更新每次写入 d 维O(n·d)线性增长KV Cache 显存n 个缓存项O(n·d)线性增长可以看到时间上的平方增长是注意力成本高的核心原因。2.3 KV Cache 同样要占用大量显存全注意力贵的第二个维度是显存。为了避免每生成一个 token 都重新计算整个历史输入的 K 和 V现代推理框架会把 K 和 V 保存下来这就是 KV Cache。KV Cache 的显存占用近似为memory 2 × num_layers × num_heads × head_dim × seq_len × bytes_per_elem这里的 2 表示 K 和 V 各有一份。假设一个 32 层模型每层 32 个头每个头维度 128使用半精度存储2 × 32 × 32 × 128 × 2 524288 bytes 512 KB / token也就是每处理一个 token大约增加 512 KB 显存。读者可能觉得单个 token 不算大但如果上下文是 10 万 tokenKV Cache 就需要约 50 GB 显存。这还没有计算中间激活值和模型参数。所以全注意力“贵”不只是慢还包括显存暴涨。上下文越长显存压力越大这就是为什么长文本任务不能无限扩展上下文的原因之一。3. 用代码验证全注意力的成本3.1 用 PyTorch 手工实现全注意力先用 PyTorch 实现一个最基础的全注意力计算函数方便观察它在不同序列长度下的表现。# 文件路径full_attention_demo.py import torch import torch.nn.functional as F def full_attention(q, k, v): q, k, v 形状均为 (batch, seq_len, head_dim) 返回经过全注意力加权后的输出和注意力分数矩阵。 d_k q.shape[-1] scores torch.matmul(q, k.transpose(-2, -1)) / (d_k ** 0.5) attn_weights F.softmax(scores, dim-1) output torch.matmul(attn_weights, v) return output, attn_weights if __name__ __main__: batch, seq_len, d 1, 8, 16 q torch.randn(batch, seq_len, d) k torch.randn(batch, seq_len, d) v torch.randn(batch, seq_len, d) out, attn full_attention(q, k, v) print(output shape:, out.shape) print(attention shape:, attn.shape)运行后可以看到输出形状为(1, 8, 16)注意力权重矩阵形状为(1, 8, 8)。这个seq_len × seq_len的二维矩阵就是全注意力中“全”字的直接体现。3.2 测量不同序列长度下的耗时下面写一个简单的测试脚本测量序列长度从 64 到 16384 时全注意力计算的平均耗时。# 文件路径benchmark_attention.py import torch import time def full_attention(q, k, v): d_k q.shape[-1] scores torch.matmul(q, k.transpose(-2, -1)) / (d_k ** 0.5) attn_weights torch.softmax(scores, dim-1) output torch.matmul(attn_weights, v) return output def benchmark(seq_len, d64, repeat20): torch.manual_seed(0) q torch.randn(1, seq_len, d) k torch.randn(1, seq_len, d) v torch.randn(1, seq_len, d) # 预热避免首次调用影响统计 for _ in range(3): full_attention(q, k, v) elapsed [] for _ in range(repeat): start time.perf_counter() _ full_attention(q, k, v) end time.perf_counter() elapsed.append(end - start) return sum(elapsed) / len(elapsed) if __name__ __main__: for n in [64, 256, 1024, 4096, 16384]: avg_time benchmark(n) print(fseq_len{n:6d}, avg_time{avg_time:.6f}s)在 CPU 环境运行输出可能接近seq_len 64, avg_time0.0012s seq_len 256, avg_time0.0031s seq_len 1024, avg_time0.0048s seq_len 4096, avg_time0.0084s seq_len 16384, avg_time0.0950s不同机器差异很大但趋势是明确的序列长度扩大 4 倍耗时可能扩大超过 4 倍。当 n 足够大时n² 增长会彻底盖过线性部分。3.3 估算 KV Cache 显存再写一个简单的脚本估算长文本场景下 KV Cache 的显存占用。# 文件路径kv_cache_memory.py def estimate_kv_cache_gb( seq_len, num_layers32, num_heads32, head_dim128, bytes_per_elem2 ): # 每个 token 需要保存的 K 和 V 元素数 elements_per_token 2 * num_layers * num_heads * head_dim bytes_per_token elements_per_token * bytes_per_elem total_bytes bytes_per_token * seq_len total_gb total_bytes / (1024 ** 3) return total_gb for seq_len in [10000, 100000, 1000000]: gb estimate_kv_cache_gb(seq_len) print(fseq_len{seq_len:9,}, kv_cache≈{gb:.2f}GB)预期输出seq_len 10000, kv_cache≈5.00GB seq_len 100000, kv_cache≈50.00GB seq_len 1000000, kv_cache≈500.00GB这只是 KV Cache 本身还没算模型参数和激活值。看到这个数字就能理解为什么长上下文推理在当前硬件上非常吃紧。4. 业界常用的缓解手段既然全注意力这么贵工程上自然不会坐以待毙。目前已经有不少优化手段大致可以分为几类。4.1 KV Cache 量化与存储优化KV Cache 不一定要用高精度存储。很多推理框架会把 KV Cache 量化成 INT8 甚至 INT4以压缩显存占用。代价是精度损失可能让输出质量下降所以如今很多框架支持按层配置量化精度。此外还有 KV Cache 复用、批处理共享等技巧。多轮对话中历史对话的 KV 可能被多个请求复用减少了重复预填充。4.2 分组查询注意力与多查询注意力标准多头注意力中每组 Q、K、V 都有独立的头。GQA分组查询注意力让多个 Q 头共享一部分 K、V 头显著减少 KV Cache 和计算量。具体来说MQA所有 Q 头共享一组 K、V GQA多个 Q 头共享一组 K、V但组数大于 1这相当于把“每页都做完整笔记”变成“多个人共用一份摘要”。代价是注意力表达的多样性可能下降但实践中效果通常可控。4.3 稀疏注意力与滑动窗口稀疏注意力不计算所有 token 对之间的关系而是只计算一部分。常见模式有局部窗口每个 token 只跟最近 w 个 token 计算注意力。全局锚点抽出一部分特殊 token与所有 token 计算注意力。随机稀疏按某种随机策略选择部分 token 对。滑动窗口是局部注意力的典型实现它把复杂度从 O(n²) 降到了 O(n·w)。问题是长程依赖可能丢失所以很多模型会混合局部窗口和少量全局 token。4.4 线性注意力路线线性注意力的核心目标是直接用数学变换把复杂度从 O(n²) 降到 O(n)。其中包括用核函数近似 softmax、将注意力转化为状态递推、使用线性 RNN 结构等方向。“Kimi Linear”这个名字如果从技术方向去理解大概率就是沿着线性复杂度这条路减少全注意力开销。关于具体的实现细节与发布口径应该以官方技术资料为准不在本文武断推断。5. Kimi Linear 与线性注意力可能的技术方向5.1 Kimi Linear 在讨论什么从最近的技术讨论看“Kimi Linear”会被和“线性注意力”“长上下文成本下降”放在一起。它的价值在于如果模型能把注意力复杂度从平方级降到线性级那么长文本场景的成本就会明显下降。不过这里要特别说明关于 Kimi Linear 的官方细节我目前没有充分掌握因此本文不编造具体参数、架构或实验数据。我们可以把它作为一个引子重点看线性注意力路线在数学上是怎么解决全注意力“贵”这个问题的。5.2 线性注意力的核心思路标准注意力公式是output_i softmax(q_i K^T) Vsoftmax 中带有指数运算无法轻易拆开成“先求和再查表”的形式。线性注意力希望把注意力近似为output_i ≈ φ(q_i) · Σ_j φ(k_j)^T v_j其中 φ 是核函数映射。这样内部先算出一个“全局状态”S Σ_j φ(k_j)^T v_j然后每次生成新 token 时只需要用当前 q 去查这个状态 S不需要再和所有历史 k、v 一一计算。这就是线性复杂度的核心从“每次重翻百万页记录”变成“维护一张不断更新的摘要卡片”。用最简单的 Python 伪代码来表达state 0 outputs [] for q, k, v in zip(Qs, Ks, Vs): state k_tensor v_tensor # 更新全局状态 output q_tensor state # 用状态生成输出 outputs.append(output)这里的内层操作跟序列长度无关只跟维度有关所以总复杂度约为 O(n·d²)而不是 O(n²·d)。5.3 线性注意力的代价线性注意力并非没有短板。softmax 被近似后注意力分布不再保证概率和为 1需要额外归一化或者设计特殊的核函数。全局状态是压缩过的长距离细粒度信息可能被折叠难以像全注意力那样精确保留每一位历史信息。对某些需要“精确匹配”的任务例如检索某一句话、精确记忆某个数字线性注意力可能不如全注意力稳定。所以线性注意力更适合上下文很长、对整体语义要求高、但对“精确逐条回溯”要求略低的场景。这也是为什么很多方案会采用“局部全注意力 全局线性注意力”的混合设计。6. 常见问题与误区6.1 问题速查表常见误解实际情况正确理解全注意力是 O(n²) 空间时间确实是 O(n²)空间还取决于是否保存注意力矩阵KV Cache 越大越好大意味着显存压力高要在缓存效率和显存成本之间取舍线性注意力完全等价于全注意力数学近似不等价存在表达能力损失需结合场景选择Kimi Linear 一定是线性注意力可能与线性复杂度方向有关具体以官方资料为准生成阶段每个 token 都重新计算所有 K/V计算量是重算注意力关系K/V 可缓存但 Q 与历史 K 的匹配仍要计算6.2 为什么预训练没有感觉那么慢有人会问预训练时上下文不也很长吗为什么大家没有集中抱怨原因在于预训练有两条保全措施并行度高所有 token 一起计算硬件利用率高。预训练对实时延迟不敏感慢一点可以接受。推理阶段不同每个新 token 都直接决定用户的等待时间。而且长对话应用中上下文可能持续累积再配合显存限制O(n²) 的代价就被格外放大。6.3 线性复杂度是不是就完全免费不是。线性复杂度的优势主要体现在“长序列”和“解码阶段”。如果序列很短线性注意力可能没优势反而因为引入额外映射或状态更复杂带来更多计算开销。具体选型要结合序列长度、任务类型和显存环境。6.4 如何判断一个模型是否适合长文本场景建议从三个维度观察复杂度机制是否引入了 KV Cache 压缩、滑动窗口、稀疏注意力或线性注意力。有效上下文模型声称支持的长度是否真的能在中长文本中保持一致性能。显存需求同等上下文长度下KV Cache 占用是否合理。7. 实战建议与深入学习路线7.1 工程上的落地建议在实际项目里减少全注意力成本可以从下面几个方向入手优先使用成熟推理框架不要自己裸写注意力实现。开启 KV Cache 量化尤其适合长对话应用。对超长文档考虑检索增强而不是让模型直接吃下全文。如果模型支持 GQA 或稀疏注意力优先选择这些版本。做性能压测时至少覆盖 1k、10k、50k 三种序列长度观察耗时和显存增长曲线。7.2 学习路径推荐如果你想把背后的原理吃透建议按下面的顺序学读懂 Transformer 原始论文中 Attention 的公式推导。手推注意力矩阵的每个维度变化。理解自回归解码与 KV Cache 的关系。阅读线性注意力的代表作包括核函数近似、线性 RNN、状态空间模型等。最后再回到 Kimi Linear 这类综合方案看它如何把这套理论工程化落地。这样一层层下来你就不会只停留在“线性比全注意力快”的表层结论而是能理解各自的利弊权衡。把“全注意力为什么贵”这个基础问题理解扎实再看线性注意力方案时会清晰很多。下一篇内容可以继续沿着注意力机制优化的方向深入拆解线性注意力是怎么用状态递推替代逐页翻阅的。
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进