全注意力机制是一切大模型的基础,也是长文本场景下成本最高的部分。最近 Kimi 团队围绕线性注意力提出了 Kimi Linear 方案,核心就是解决全注意力“越用越贵”的问题。这篇作为“核心原理”系列的第一篇,先把全注意力为什么贵这件事讲透:从计算复杂度、KV Cache、生成阶段的重复开销三个角度拆解,再看 Kimi Linear 想动手术的位置在哪里。
如果你在纠结“为什么模型处理长文本时每生成一个新 token 都越来越慢”,或者想理解线性注意力到底优化了什么,这篇文章可以直接往下看。本文不涉及具体部署命令,重点做原理拆解和成本量化,适合算法工程师、LLM 应用开发者以及对推理优化感兴趣的技术读者。
1. 核心原理速览
| 维度 | 说明 |
|---|---|
| 主题 | 标准全注意力机制的计算开销来源 |
| 核心复杂度 | 序列长度 N 下,注意力矩阵计算与显存均为 O(N²) |
| 关键瓶颈 | QK^T 矩阵、Softmax 归一化、KV Cache 重复读取 |
| 生成阶段特点 | 每个新 token 都要与全部历史 token 计算注意力 |
| 典型影响 | 长上下文下推理时延增长、显存占用上升 |
| 优化方向 | FlashAttention、稀疏注意力、线性注意力、Kimi Linear |
| 适合读者 | 想理解 LLM 长文本成本来源的开发者与算法工程师 |
| 前置知识 | 熟悉 Transformer 基础结构、了解 QKV 概念 |
这里说明一下:本文讨论的是全注意力机制在推理阶段(尤其是自回归生成阶段)的成本模型。Kimi Linear 的具体实现细节目前以官方技术报告为准,本文只从公开原理层面说明它要解决的问题,不虚构参数。
2. 直观理解:为什么叫“重翻百万页记录”
把大模型生成文本想象成一个非常认真的抄写员。它在写当前这句话时,每写一个词,都要把之前所有写过的内容重新看一遍,确认这个词和前面每个词的关系。这个“重新看”的动作就是全注意力。
如果文章只有一句话,重新看一遍很快;如果文章有一百万字,那么每写一个新词,就要翻一百万字的历史记录。而且麻烦的是:不是只看一遍,而是每个词都要翻一遍。所以生成第 10 个词时,它翻 9 条记录;生成第 100 个词时,它翻 99 条记录;生成第 10000 个词时,它翻 9999 条记录。整个过程累加出来,就是 O(N²) 的开销。
这个比喻基本准确。在实际实现中,“翻记录”并不是真的重新读一遍原始文本,而是读取缓存下来的 K 矩阵和 V 矩阵。K 和 V 分别代表历史 token 的“检索键”和“内容值”,它们在推理时被缓存在显存里,统称 KV Cache。每个新 token 的 Q 向量要跟所有历史 K 向量做点积,再把结果和所有 V 向量加权求和。
也就是说,即使没有重新计算前面的隐藏层状态,注意力层的开销依然随上下文长度线性增长,而这个线性增长发生在每个生成步骤里,最终形成 O(N²) 的总成本。
3. 从公式看成本:全注意力的时间复杂度
标准注意力公式为:
Attention(Q, K, V) = softmax(QK^T / sqrt(d_k)) V其中 Q、K、V 的形状都是 (N, d),N 是序列长度,d 是每个头的维度。这里有一个非常关键的操作:QK^T。
QK^T 的矩阵乘法是把 Q 的每一行和 K 的每一行做点积,得到一个形状为 (N, N) 的注意力分数矩阵。这个矩阵就是所有 token 两两之间的关联程度。问题在于:它的规模是 N²。
如果 N 是 1000,矩阵是 100 万个数;如果 N 是 10 万,矩阵是 100 亿个数;如果 N 是 100 万,矩阵就是一万亿个数。这个增长速度非常夸张,远快于显存和算力本身的增长。
下面用一段 Python 代码模拟注意力矩阵的显存占用:
import math def attention_matrix_bytes(seq_len, head_dim, dtype_bytes=2): # QK^T 矩阵元素个数为 seq_len * seq_len elements = seq_len * seq_len # 显存占用还要乘以每个元素占用的字节数 memory = elements * dtype_bytes return memory for n in [1000, 10000, 100000, 1000000]: mem = attention_matrix_bytes(n, 128) print(f"序列长度 {n:>9,} -> 注意力矩阵显存占用 {mem / 1e9:.2f} GB (fp16)")输出大致为:
序列长度 1,000 -> 注意力矩阵显存占用 0.002 GB 序列长度 10,000 -> 注意力矩阵显存占用 0.20 GB 序列长度 100,000 -> 注意力矩阵显存占用 20.00 GB 序列长度 1,000,000 -> 注意力矩阵显存占用 2000.00 GB这只是单个注意力头的 QK^T 矩阵。实际模型有多个头、多层,还要再乘以头数和层数。所以当上下文来到十万甚至百万 token 这个量级,全注意力在显存和时间上的开销都会变得不可接受。
4. 训练阶段 vs 推理阶段:成本的两种形态
全注意力的昂贵,在训练和推理阶段表现不一样,不能混在一起谈。
训练阶段是并行计算整个序列的。也就是说,模型一次读入所有 token,一次性算出所有 token 两两之间的注意力分数。此时成本主要体现在显存上:N² 的注意力矩阵必须被实例化(或者通过 FlashAttention 这样的内核融合技术避免完整实例化),同时反向传播还要保存中间结果。因此训练长序列模型的最大瓶颈是显存。
推理阶段又分为两个子阶段:
预填充阶段(Prefill):模型拿到用户输入的一整段文本,并行计算所有输入 token 的注意力。此刻成本形态和训练类似,但不需要反向传播,压力比训练小很多。
自回归生成阶段(Decode):模型逐 token 生成输出。每生成一个 token,都要让这个新 token 的 Q 和所有历史 token 的 K、V 做注意力计算。这里的问题不在于单次计算量有多大,而在于它被重复执行 N 次。序列越长,单次需要读取的 KV Cache 就越大,计算延迟随之增加。
很多人在实际使用长文本模型时感觉到“越到后面越慢”,就是第二个阶段的问题。它并不是心理作用,而是每一次生成新 token 都需要访问越来越大的 KV Cache。
5. 生成阶段的真实瓶颈:KV Cache 的重复读取
继续前面抄写员的比喻。“重翻百万页记录”在工程上的真实形态,是每一层、每一步都要读取完整的 KV Cache。
假设一个模型有 32 层、40 个注意力头、每个头维度 128,上下文长度为 100 万 token,使用 fp16 存储,KV Cache 的大小可以估算:
def kv_cache_bytes(layers, num_heads, head_dim, seq_len, dtype_bytes=2): # K 和 V 各一份 per_layer = 2 * num_heads * head_dim * seq_len * dtype_bytes return per_layer * layers layers = 32 num_heads = 40 head_dim = 128 seq_len = 1000000 total = kv_cache_bytes(layers, num_heads, head_dim, seq_len) print(f"总 KV Cache 大小: {total / 1e12:.2f} TB") print(f"平均每层: {total / layers / 1e9:.2f} GB")这个量级已经远超单张显卡的显存容量。就算不考虑显存放不放得下,在生成每个新 token 时,要把这么大体量的 K、V 数据从 HBM 里读取出来做矩阵乘,这个内存访问开销本身就会成为延迟的主要来源。
换句话说,全注意力在生成长文时的瓶颈不只是“计算量大”,还有“每个 token 都要把所有历史数据重新读一遍”的访存开销。这也是为什么很多长文本优化方案都在做 KV Cache 压缩、剪枝、滑动窗口、或把注意力变成线性形式,本质上都是在减少生成阶段需要重复读取的数据量。
6. 已有优化路线:FlashAttention、稀疏注意力、线性注意力
在讨论 Kimi Linear 之前,有必要把既有的注意力优化路线梳理一遍。它们解决问题的角度各不相同。
FlashAttention 属于“把计算重排”的路线。它不改变注意力的数学定义,而是通过分块计算、内核融合,避免把完整的 N×N 注意力矩阵写入全局显存。这样可以在同样显存下处理更长的序列,同时减少显存读写,训练速度也能提升。但 FlashAttention 并没有把复杂度从 O(N²) 变成 O(N),它只是把 N² 的显存压力通过分块技术缓解了一部分,计算量依然是 N²。
稀疏注意力属于“减少计算范围”的路线。它假设不是所有 token 都同等重要,用一个固定模式限制每个 token 只能关注部分历史 token。典型做法有滑动窗口注意力、全局锚点 token 加局部窗口等。好处是复杂度可以降为 O(N),坏处是模型的表达能力受限,某些需要跨长距离关联信息的任务可能效果下降。
线性注意力属于“改变计算顺序”的路线。传统注意力必须先算 QK^T 得到 N×N 的注意力矩阵,再和 V 相乘。线性注意力通过矩阵乘法的结合律,调整计算顺序,把对 N² 矩阵的需求变成 N 量级。它维护一个全局的状态矩阵,基于这个状态逐步更新结果。理论复杂度 O(N),但早期线性注意力在效果上往往不如标准注意力。
Kimi Linear 本质上属于线性注意力路线的探索。它要解决的核心问题,就是如何让线性注意力既保持标准注意力的表达能力和实际效果,又把生成阶段的成本降下来。
7. Kimi Linear 要解决什么
Kimi Linear 这个名字本身已经说明了方向:用线性复杂度的注意力替代二次复杂度的全注意力。结合上面分析,它想解决的是全注意力在超长上下文下的两个核心问题:
第一,生成阶段的 KV Cache 过大。如果注意力计算方式改为线性,就不再需要缓存完整的 K 和 V 矩阵,而是维护一个固定大小的状态。这个状态大小可以做到与序列长度无关或弱相关,显存占用从 O(N) 降到 O(1) 或 O(log N) 级别。
第二,每生成一个新 token 就要重读全部历史记录的访存问题。线性注意力把历史信息压缩成一个固定大小的状态,生成新 token 时只需要读取这个状态,而不需要读取全部历史 K、V。时间开销从 O(N) 降到 O(1)。
但是这里要提醒一下:线性注意力不是没有代价。传统全注意力的每一个 token 都可以直接访问所有历史 token 的精确向量,信息检索能力很强。线性注意力把历史信息压缩成固定大小状态后,信息存储容量受限,这可能导致模型在需要精确回忆某些细节的任务上表现下降。Kimi Linear 的技术核心,大概率就是在解决表达能力和复杂度之间的平衡问题。
具体实现细节需要以 Kimi 团队发布的技术报告为准。但从原理层面可以确定的是,如果线性注意力能够在长文本任务上达到接近全注意力的效果,那么它对超长上下文的推理成本改善会是数量级的。
8. 如何量化观察:复杂度模拟与实测建议
原理讲完了,接下来给出一个可操作的量化思路。虽然这里不涉及具体模型部署,但你可以用下面方法观察项目里的注意力成本。
8.1 FLOPs 理论估算
全注意力的计算量可以用公式快速估算。单层单头的 QK^T 计算量约为 2 × N² × d,乘以 V 的计算量也有类似量级。写一个简单模拟:
import math def attention_flops(seq_len, head_dim, layers, num_heads): # 单头 QK^T: 2 * N * N * d # 单头 attn @ V: 2 * N * N * d per_head_flops = 2 * seq_len * seq_len * head_dim * 2 total_layers = layers * num_heads * per_head_flops return total_layers for n in [10000, 50000, 100000, 500000]: flops = attention_flops(n, 128, 32, 40) print(f"seq_len={n:>7,} -> 注意力总计算量约 {flops / 1e15:.2f} PFLOPS")这个数量级可以让你直观理解为什么长序列下全注意力很难跑起来。实际部署时,真实的计算时间还取决于硬件算力和内存带宽,但这个理论值能帮助判断瓶颈在哪个环节。
8.2 实测观察指标
如果你已经在本地或服务器上跑某个 Transformer 模型,建议观察以下指标:
- 每生成一个 token 的耗时(decode time per token)
- KV Cache 占用多少显存
- 模型整体显存占用随上下文长度的增长曲线
- 长文本输入时,预填充阶段和生成阶段的耗时分布
# 观察显存占用,每秒刷新一次 nvidia-smi --query-gpu=memory.used,memory.total,utilization.gpu --format=csv -l 1如果发现显存占用随输入长度线性上升,而生成耗时也明显上升,说明瓶颈在 KV Cache 访存和注意力计算。此时再去考虑切换到稀疏注意力或线性注意力方案才有依据。
8.3 对比测试方法
要验证某个优化方案是否有用,最直接的方式是设置两个实验组:一组使用标准全注意力,另一组使用优化后的注意力(如线性注意力、稀疏注意力)。固定相同模型结构、相同输入数据、相同 batch size,分别测量:
- 峰值显存
- 每秒生成的 token 数
- 相同 prompt 下的输出质量
- 长文本场景的任务准确率
这里要特别强调:不能只看速度,还要看质量。很多线性注意力方案在短文本上速度提升不明显,在长文本上才能体现优势,但文本质量可能有所下降。所以测试时文本长度要覆盖短、中、长三档,不要只看某一档。
9. 常见误区与排查思路
全注意力昂贵这个结论本身很清晰,但实际工程里有一些常见误区值得澄清。
| 误区 | 实际情况 | 建议 |
|---|---|---|
| 长文本慢是因为模型参数量大 | 参数量不变时,注意力开销随序列长度平方增长是核心因素 | 先看长度,再谈参数量 |
| FlashAttention 把复杂度变成了 O(N) | FlashAttention 只是减少了显存读写和内存占用,计算量仍是 O(N²) | 超长上下文仅靠 FlashAttention 不够 |
| 稀疏注意力一定比全注意力好 | 固定稀疏模式可能丢失关键远距离信息 | 根据任务类型评估效果 |
| 线性注意力一定能无损替代全注意力 | 信息压缩会带来表达上限,质量和速度需要平衡 | 验证具体任务指标 |
| 显存不够就加个显卡 | 多卡还要考虑通信开销,KV Cache 分发也有成本 | 先评估优化注意力本身 |
如果你在自己的项目里遇到“长文本推理越来越慢”的问题,排查顺序建议是:
- 确认序列长度是否是主要变量:把长度减半,看耗时是否明显下降。
- 确认是预填充阶段慢还是生成阶段慢:分别统计两个阶段的耗时。
- 确认 KV Cache 是否被完整加载:用 profiler 看 attention 算子的耗时占比。
- 确认是否已经启用 FlashAttention 或类似内核融合优化:很多框架默认不开启,需要显式设置。
- 确认是否有不必要的中间张量被保存:推理模式下关闭梯度、关闭中间状态保存。
10. 最佳实践与使用建议
针对全注意力成本问题,下面给出几条工程建议。
第一,短文本场景不要盲目换线性注意力。如果上下文长度只有几千 token,全注意力的计算开销并不高,换成线性注意力反而可能因为表达能力下降而影响效果。优化手段要匹配实际瓶颈。
第二,长文本场景先量化再优化。不要凭感觉判断“慢是因为注意力”。先用 profiler 和 nvidia-smi 定位瓶颈,确认注意力算子确实占用主要耗时后,再考虑稀疏化或线性化。
第三,关注生成阶段多于预填充阶段。对于对话、写作、代码生成等交互式应用,用户感受到的延迟主要来自逐 token 生成。预填充阶段虽然计算量大但只跑一次,生成阶段则要跑 N 次,优化价值更高。
第四,注意输出质量的回归测试。任何注意力优化方案都可能改变模型行为。建议准备一套覆盖摘要、推理、代码、多轮对话等任务的小型评测集,在切换注意力实现后跑一遍对比,防止“速度上去了,效果下来了”。
第五,结合上下文工程来降低实际序列长度。不是所有场景都需要模型处理百万 token。检索增强、分块摘要、滑动上下文等方案都能显著降低注意力开销,而且效果稳定、风险低。注意力层面的优化可以作为进一步的性能手段,而不是唯一解法。
11. 总结与下一步
这篇文章把全注意力贵在哪里讲清楚了:核心是 O(N²) 的注意力矩阵计算,加上生成阶段每个 token 都要反复读取 KV Cache,二者叠加,导致长文本推理的时间和显存开销随长度急剧增长。Kimi Linear 瞄准的正是这两个瓶颈,通过线性注意力的方式把复杂度从 N² 拉低到接近线性。
读完这篇,你应该先做一件事:用文中的公式估算一下你实际场景里的注意力开销,再判断瓶颈到底在计算量还是访存。不要急着换模型,用数据说话。
下一篇“核心原理 02”可以继续拆解 Kimi Linear 的技术细节:它是如何压缩历史信息的、状态是怎么维护的、跟现有线性注意力有什么差异。如果你在实践里遇到了长文本推理延迟和显存问题,建议先收藏本文,后面排查时可以对着复杂度模型逐项定位。