ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

全注意力为何昂贵?线性注意力如何用固定状态压平复杂度

全注意力为何昂贵?线性注意力如何用固定状态压平复杂度 全注意力为什么贵最直观的场景就是长文本生成。Kimi Linear 这类面向超长上下文的架构要解决的第一道问题就是 Attention 的平方复杂度。这篇是 Kimi Linear 核心原理系列的第一篇先把最底层的账算清楚模型每生成一个新 token为什么要把已有的上下文全部重新读一遍这遍“重读”到底消耗了什么线性注意力又是用什么思路把消耗降下来。适合已经知道 Transformer 里有 QKV、但对复杂度成本没有完整账本的开发者阅读。读完这篇你能解释 O(n²) 从哪来、KV Cache 解决了什么没解决什么、以及 Linear Attention 如何用一个固定大小的状态换掉整份 KV 历史。1. 一个直观比喻生成新词相当于带着问题重查整本档案1.1 注意力机制本质上是一次“按相关度加权”的档案检索先建立比喻。Transformer 的 self-attention 里当前 token 会生成一个 Query查询条件历史每个 token 各生成一个 Key档案标签和一个 Value档案正文。新 token 的输出等于把历史所有 Value 按照“Query 与 Key 的相似度”做加权求和。相似度越高权重越大对应档案正文对输出的影响越大。用代码表达这一步就是下面这个循环def attention_for_one_token(q, keys, values): scores [] for key in keys: scores.append(dot(q, key)) weights softmax(scores) out zeros_like(values[0]) for w, v in zip(weights, values): out w * v return out这个循环就是全注意力的最小单元一个 query对全部 key 做点积再对全部 value 做加权求和。key/value 列表每增加一个这个循环就多一轮。当序列长度为 n 时n 个 query 都要面对 n 个 key所以整体规模是 n×n 量级。这里有一个很容易忽略的细节softmax 的权重必须等所有分数算完之后才能归一化。也就是说单个 query 的输出依赖整段历史的全部 key无法边读边输出必须先完整扫描一遍。这正是“重翻档案”这个比喻的来源之一。1.2 “每个新词都要重翻百万页”到底指什么在因果语言模型里生成第 n 个 token 时它只能看到前 n-1 个 token。要计算这个 token 对所有历史 token 的注意力就必须把前 n-1 个 key 全部取出来算相似度。生成第 n1 个 token 时key 的数量变成 n。每前进一步要处理的档案就多一页。从头到尾生成 N 个 token总共处理的 key-value 数量近似是 1 2 ... N也就是 N²/2 量级。所以上下文长度翻一倍生成成本大约变四倍。上下文长度提高 10 倍生成成本大约提高 100 倍。标题里“重翻百万页记录”说的不是某一次操作特别贵而是每一步都在递增地重翻。生成第 10 万个 token 时它要把前 99999 个 token 的 key 全部扫一遍生成第 20 万个 token 时扫描量又翻一倍。累积下来这就是平方成本最直觉的体现。2. 全注意力的成本模型算力、显存和内存带宽三笔账2.1 训练和预填充阶段n×n 分数矩阵直接爆炸在训练和预填充阶段所有位置可以并行计算注意力这时Q K^T会生成一个 n×n 的分数矩阵任意两个位置都要配对。计算量是O(n²·d)显存开销是O(n²)。当上下文长度 n 达到 10 万时分数矩阵是 10^10 个元素即使只按 float16 计算也要 200 GB 以上显存这还没算模型参数和中间激活。正式因为这样全注意力在超长上下文的预填充阶段必须依赖分块、稀疏化或 Linear 方案否则一张卡根本放不下。2.2 自回归解码阶段每一步都重读一遍全部 KV推理生成时模型通常维护 KV Cache。每个历史 token 的 Key、Value 已经算好放在缓存里新 token 只需要算自己的 Query再和缓存里的所有 Key 做点积。看起来很省但每一步仍然要读取全部缓存。这一阶段真正的瓶颈往往不是算力而是内存带宽。假设上下文已经有 10 万 token每层 KV 缓存占几百兆字节每生成一个 token 就要把这些数据从显存读一遍。生成一百个 token就要读一百遍。长上下文对话里那种“开头很快、越说越慢”的感觉主要就是读取量随上下文长度线性增长造成的。用表格汇总全注意力的成本阶段全注意力成本主要瓶颈训练/预填充O(n²) 计算、O(n²) 显存n×n 分数矩阵不可承受解码单步O(n) 计算、O(n) 内存读取KV 缓存读取带宽完整生成 N 个 tokenO(N²)每一步的 n 都在增长KV 缓存空间O(N·L·H·d)长上下文显存占用2.3 KV Cache 解决了什么没解决什么KV Cache 解决的是“重复计算”没有缓存时生成第 N 个 token 要把前 N-1 个 token 的 K/V 全部重算一遍有缓存后每一个历史位置的 K/V 只算一次。但它没有解决“重复读取”每个新 token 仍然要把缓存里的所有 Key 拿出来做点积把所有 Value 拿出来做加权。这笔账要分开算缓存空间随上下文线性增长。单步读取量随上下文线性增长。生成全程的总读取量是平方增长。所以 KV Cache 不是长上下文的万能钥匙。它让“不重算”成为可能却让“每一步都重读”成为生成阶段的主导成本。这也正好对标题里的比喻不是翻一次档案而是每写一页新记录就把前面所有记录重新翻一遍。3. 线性注意力的核心替换把“逐页重读”改成“维护摘要”3.1 为什么 softmax 做不到增量更新先看 softmax 注意力的问题。第 i 个位置的输出写成out_i Σ_j softmax(q_i · k_j) · v_jsoftmax 的分母是Σ_j exp(q_i · k_j)它依赖全部 key。这意味着你无法只保留一个固定大小的状态然后准确算出这个归一化项。只要相似度函数还是 softmax历史信息就必须以完整的 key/value 形式保留供未来每个 query 重新计算。想真正把复杂度降下来不能只在 softmax 外面做工程优化而是必须换算子本身。这是线性注意力和全注意力最本质的分界线。3.2 换一种聚合顺序让固定状态成为可能线性注意力的做法是把 softmax 相似度替换成两个特征向量φ(q)和φ(k)的点积再利用矩阵乘法结合律调整计算顺序。原来的公式是“先算Q K^T再乘 V”线性注意力把它改成“先算K^T V再乘 Q”全注意力: Output softmax(Q K^T) V 线性注意力: Output φ(Q) (φ(K)^T V)这里的φ(K)^T V是一个形状固定的小矩阵可以随着 token 逐个到达不断累加。计算过程变成def linear_attention_step(state, z, q, k, v): # state: (d, d) 历史累计摘要 # z: (d,) 历史 key 特征累计 phi_q feature_map(q) # (d,) phi_k feature_map(k) # (d,) state outer(phi_k, v) # 新档案进入摘要 z phi_k out phi_q state / (phi_q z) return out, state, z线性注意力在生成第 n 个 token 时不需要保存前 n-1 个 key 和 value只需要保存累计状态 S 和累计归一化向量 z。新 token 到达后把φ(k) × v累加到 S把φ(k)累加到 z当前 token 的输出用φ(q)^T S除以φ(q)^T z得到。对应到标题的比喻不再是“每写一页新记录就重翻整本档案”而是“每来一页新记录就更新一页摘要查询时只读摘要”。3.3 固定状态意味着信息压缩也意味着表达力边界状态矩阵 S 的大小由d×d决定d 是特征映射维度d 是模型维度。两者都不随上下文长度变化。线性注意力把“全文检索”换成了“固定大小摘要”。好处成本线性化状态固定延迟稳定。代价信息有损。摘要能保留统计规律和整体模式但很难精确还原某一条历史记录。如果任务要求“从第 5000 段摘录某个精确数字”线性注意力通常比全注意力吃力。这个“有损”不是实现 bug而是设计选择。选择哪种注意力本质是在“长效记忆的压缩程度”和“精确检索能力”之间找平衡点。4. Kimi Linear 这一类方案为什么能撑起超长上下文4.1 成本曲线从平方变线性长度越大差距越大用数量级对比看上下文长度全注意力总体操作数约线性注意力总体操作数约1k10^610^3 × d²100k10^1010^5 × d²1M10^1210^6 × d²同一份上下文长度从 10 万提高到 100 万全注意力的成本要放大 100 倍线性注意力只放大 10 倍。这才让“百万 token 上下文”有了工程可行性也让 Kimi Linear 这类名字里的“Linear”有了实际意义它强调的不是某一种函数而是把长上下文处理从平方复杂度压到线性复杂度。需要说明的是线性注意力并不等于完全取消 KV 存储而是把注意力内部的历史依赖压缩成状态。具体实现里状态维度、特征函数怎么选、是否需要组合滑动窗口是系列后续内容要展开的细节这一篇先把成本原理立住。4.2 工程实现里的数值稳定、归一化和特征映射问题替换 softmax 之后会冒出三类工程问题特征映射选择。常见有φ(x)elu(x)1、ReLU、随机特征近似等不同的 φ 影响表达力和数值范围。归一化维护。分母φ(q)^T z必须维护否则输出尺度会随着历史变长越来越大。长期累加精度。状态 S 一直在累加float16 下可能逐渐积累舍入误差需要定期检查或者改用分块计算。生产环境通常要配合分块计算chunked attention把长序列切成块块内用相对高精度的方式聚合块间做增量更新。这样既能保留并行度又能控制数值误差。这个主题值得单独开一篇细讲。4.3 混合架构局部精确检索加全局摘要Kimi Linear 这类方案能撑长上下文不等于所有层、所有位置都该改用线性注意力。实际系统里更多是混合方案局部窗口内用精确 attention保住细节检索能力。全局范围用 linear state拿到低成本的长距离依赖。或者前几层保持全注意力深层才使用线性注意力。选型时可以问自己三个问题应用里是否需要精确回忆某条历史信息上下文是否真的会长期增长到万级、十万级以上当前瓶颈到底是训练显存、推理延迟还是 KV 缓存空间回答完这三个问题再决定该用全注意力、线性注意力还是两者混合。5. 训练与推理的工程差异以及上线前要补的账5.1 预填充阶段序列长度才是显存压力来源预填充要一次性处理整个 prompt。全注意力在这个阶段会产生很大的分数矩阵显存风险最高线性注意力则适合用分块方式在长序列上并行计算状态再合并结果。在学习和实验环境里序列长度只有几百几千时两者差异不明显。一旦把序列长度推到 8k、32k显存和延迟的差异会立刻拉开。做长上下文实验时建议从短序列开始验证正确性再逐步拉长观察显存曲线。5.2 解码阶段固定状态带来稳定延迟全注意力在解码阶段每一步的延迟都随上下文长度上涨因为要读的 KV 越来越多。线性注意力的每步延迟基本稳定因为状态大小固定。对实时对话、智能体工具调用这类延迟敏感场景这是一个非常大的差异。想象一个已经读了 50 万 token 资料的智能体每回答一个问题都要把 50 万 token 的 KV 全部扫一遍和每次都只读一个固定大小的状态体验完全不同。5.3 生产环境额外要考虑的东西注意力复杂度只是模型侧的一笔账真正上线时还要补上请求级上下文管理对话历史如何保留、截断、压缩。状态生命周期线性注意力状态什么时候初始化、什么时候重置。日志和监控单步延迟、KV 缓存占用、状态范数是否异常。精度策略fp16、bf16、fp32 混合使用时的误差控制。回滚方案新注意力实现出现数值问题时能否快速切回旧版本。这些和注意力原理无关但决定一个方案能不能稳定落地。6. 常见误区、排查路径和选型建议6.1 三个经常想错的点误区一已经有了 KV Cache就不怕长上下文了。KV Cache 解决的是重复计算没有解决重复读取。KV 总大小随 token 数线性增长每个新 token 都要重新读一遍缓存生成阶段总读取量仍然是平方的。KV Cache 是优化不是复杂度解药。误区二线性注意力就是把公式换一下顺序。它确实把“先算Q K^T再乘 V”换成了“先算K^T V再乘 Q”但这个替换只有在相似度可以被特征点积表示时才成立。softmax 本身做不到这一点所以必须换算子。这个变化不是实现细节而是表达能力的改变。误区三复杂度是 O(n)就说明成本很低。线性注意力的状态更新仍有固定开销特征映射和归一化也会引入额外计算。短上下文场景下线性注意力的优势不明显甚至可能因为压缩误差而表现更差。O(n) 说的是增长趋势不是绝对值。6.2 线上排查清单如果长上下文模型表现异常按这个顺序排查检查项检查方式说明解码延迟是否随长度线性增长记录不同上下文长度下的每 token 延迟接近线性说明读取成本可控KV 缓存占用曲线监控每请求 KV 字节数确认是否有异常膨胀状态数值是否漂移对比长序列早期与后期输出范数发现漂移优先检查归一化分母精确检索任务表现构造“抄第 N 段原文”测试评估是否需要用混合注意力数值精度是否下降比较 fp16/bf16/fp32 输出差异累加型状态对精度更敏感6.3 学习环境下的验证顺序自己动手复现时不要一开始就冲百万 token。建议按这个顺序在几千 token 的 toy 数据集上分别训练全注意力和线性注意力模型。对比 loss 曲线、收敛速度和一个简单检索任务的准确率。测量不同上下文长度下的前向耗时和显存占用画出复杂度曲线。跑通之后再尝试分块计算、状态压缩、混合注意力。这样每一步都能用可控实验定位问题也能亲眼看到平方和线性两条曲线的差距。7. 这一篇的收尾判断和后续方向7.1 一句话记住这套成本关系全注意力为什么贵每个新 token 的注意力输出都依赖全部历史 key/valuesoftmax 的分母又让历史无法被提前压缩成固定状态所以每一步都要重读整段历史。KV Cache 省了重复计算没有省重复读取真正把“重翻全书”改成“维护摘要”的是线性注意力用固定状态换掉全量 KV。7.2 系列下一篇可以展开什么这一篇是 Kimi Linear 核心原理的第一篇后续值得继续展开的主题包括特征映射的选择不同的 φ(q)、φ(k) 如何影响状态表达力和数值稳定性。分块训练chunked 方式在长序列上如何并行累加状态、控制误差。混合架构局部精确注意力与全局线性状态如何组合。百万 token 场景下的训练和推理流水线。对新手来说最有价值的练习不是先调参而是先在一个小数据集上把全注意力和线性注意力的复杂度曲线画出来。看到平方和线性在 10 倍长度上的差距之后再回头看这篇里的每一句都会更有体感。
返回列表