尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

全注意力机制为什么贵?从O(n²)到线性注意力成本优化解析

全注意力机制为什么贵?从O(n²)到线性注意力成本优化解析 如果你训练或推理过大语言模型大概率遇到过这样的现象序列长度从 2048 加到 4096你以为只是数据量翻一倍结果显存直接不够用了训练速度也肉眼可见地掉下来。如果再加到 8192、16384很多机器连试都不敢试。这不是优化不到位而是全注意力机制在结构上就决定了序列越长成本增长越夸张。它像一位极其尽责的档案管理员你每问一句新问题他不是翻最后一页而是把从第一页到最新一页的所有档案全部重新读一遍。这篇文章要讲清楚一件事为什么全注意力机制这么“贵”它的成本模型究竟是怎么拆的Kimi Linear 这类以 Linear 为名的方案为什么被认为是解决这个问题的方向我不打算停留在“O(n²) 复杂度”这个公式层面而是会把计算成本、显存成本、推理成本分开拆配合可运行的最小代码示例让你真正理解“每个新词都要重翻百万页记录”的含义。读完这篇文章你会知道全注意力为什么是二次方复杂度以及这个二次方究竟花在了哪里训练和推理阶段注意力机制的成本结构有什么不同工业界已经有哪些降低注意力成本的方案它们各自牺牲了什么从“全注意力”到“Linear Attention”这类方案核心思路发生了怎样的转变实际工程里你应该怎么选择方案怎么验证效果怎么避坑。1. 全注意力机制到底在计算什么在深入成本分析之前先把全注意力的计算过程说清楚。注意力机制Attention最早被广泛认知是因为 2017 年的 Transformer 论文Attention Is All You Need。它的核心想法是序列里的每个 token可以简单理解为一个词或一个字在生成新表示时不应该只看自己而应该参考序列里所有其他 token并根据相关性分配不同的权重。这里有一个公式也是全注意力最经典的定义Attention(Q, K, V) softmax(Q K^T / sqrt(d_k)) VQ是查询Query表示“我现在想找什么信息”K是键Key表示“我这条信息是什么主题”V是值Value表示“我这条信息的实际内容”d_k是键向量的维度用于缩放点积结果。计算分三步用Q和所有K做点积得到每个位置和其他所有位置的相似度分数对分数做 softmax归一化成权重用权重对所有V加权求和得到最终的输出。这个过程里关键在第 1 步对于序列里的每一个 token都要和其他所有 token 计算两两之间的相似度。序列长度如果是n那相似度矩阵的大小就是n × n。如果你曾经维护过一张两两配对的表就知道这里的增长有多快10 个元素有 45 对100 个元素有 4950 对1000 个元素有 499500 对。对不是 100 倍而是接近 10000 倍。这就是 O(n²) 的来源——两个嵌套循环每一个 token 都要遍历一遍全部 token。这个设计的好处是模型不会遗漏任何跨位置的信息。无论两个词在序列里隔得多远注意力都能直接建立联系。但坏处也在这它不做任何取舍永远在完整地计算两两关系。2. “重翻百万页记录”到底是什么意思你可以把全注意力理解成一个极其负责、但不太聪明的档案管理员。他的工作方式是你问他一个问题新的 token他会从档案柜里取出所有历史记录前面的全部 token把新问题和每一页档案逐一比对计算相关程度最后把相关程度最高的几页内容摘出来作为回答依据。问题在于他从来不会把比对结果缓存起来。下一次你再问一个新问题他又把整个档案柜从头翻一遍。序列越长档案越厚他每回答一个新问题消耗的时间就越长。这个比喻对应的是 Decoder 推理阶段的行为。生成第n1个 token 时模型需要将新 token 的 Query 和之前所有 token 的 Key、Value 重新计算一遍注意力。虽然 Key 和 Value 本身可以缓存这就是 KV Cache 的由来但注意力权重计算这一步仍然需要和全部历史 token 逐一交互。这就带来了两个直观后果总成本随序列长度线性增长——每生成一个新词都要翻一遍越来越厚的档案KV Cache 也随之增长——每多一个 token就要多存一组 Key 和 Value。这就是为什么长文本场景下推理引擎最怕的不是计算而是“上下文越来越长”。从表面看模型只是在多生成一个词但它的成本却取决于这个词离输入开头有多远。你让模型生成第 100 个词和第 10000 个词后者的单步成本要高得多。3. 成本拆解训练、推理和显存各贵在哪标题里说“全注意力为什么贵”这个“贵”不能笼统地说要拆成三个维度看。3.1 训练阶段计算量随序列长度二次方增长训练阶段的最大特点是序列里的所有位置都是已知的我们可以并行计算所有 token 之间的注意力。这让训练比推理快很多但也带来了一个问题所有 token 两两之间的注意力分数都要计算一遍。假设输入序列长度是n每个 token 的向量维度是d。一次注意力计算中计算Q K^T需要n × d和d × n的矩阵乘法得到n × n的分数矩阵这个矩阵的大小是n²对n²个元素做 softmax再用n²个权重去加权n × d的 V 矩阵。所以单次注意力的时间复杂度和空间复杂度都是 O(n²)。Transformer 里通常有多个注意力头假设有h个头每个头的维度是d_h那总的复杂度仍然是 O(n²)只是系数更大。直观感受一下训练一个 2K 上下文的模型单条样本的注意力分数矩阵是 2048 × 2048约 400 万个元素。如果上下文提升到 8K矩阵变成 8192 × 8192约 6700 万个元素。序列变成 4 倍计算量变成了 16 倍。这就是为什么长上下文训练那么贵的根本原因。3.2 推理阶段KV Cache 让显存压力线性增长推理阶段大家都会用 KV Cache 来做加速。所谓 KV Cache就是把已经计算过的 Key 和 Value 缓存下来避免生成每个新 token 时重新计算前面所有位置的 K 和 V。这个策略很有效但也带来了新的成本。假设模型有L层、h个注意力头、每个头的维度是d_h序列长度是n那么每一层需要缓存n × h × d_h的 Keyn × h × d_h的 Value。两者合计2 × n × h × d_h。再乘以层数LKV Cache 的总大小就是2 × L × n × h × d_h。用具体数字感受一下。一个 7B 规模的模型层数约 32 层注意力头数约 32 头每头维度约 128。上下文长度如果是 8192那么 KV Cache 大约是2 × 32 × 8192 × 32 × 128 2,147,483,648 个 float16 数值换算成显存约 4 GB。这还没算模型参数本身和中间激活值。如果上下文长度翻到 32K光 KV Cache 就要 16 GB 以上。很多显卡根本放不下。这就是为什么长文本推理时你常常看到显存被 KV Cache 吃满而不是被模型参数吃满。模型参数是固定的KV Cache 却随着对话轮数和上下文长度动态增长。3.3 模型参数量增加带来的附加成本还有一个容易忽略的成本为了提升注意力表达力模型往往需要更大维度的d和更多层数。维度越大Q、K、V 的线性投影矩阵也越大但这部分成本是线性增长的不是全文重点。真正导致长文本场景崩溃的仍然是注意力分数矩阵的二次方增长。4. 用代码理解复杂度一个最小示例纸上谈兵容易直接看代码最直观。下面用 Python 和 NumPy 实现一个简化版的全注意力计算并测量不同序列长度下的耗时变化。import numpy as np import time def full_attention(Q, K, V): # Q, K, V shape: (n, d) d_k Q.shape[-1] scores Q K.T / np.sqrt(d_k) # (n, n) weights np.exp(scores - np.max(scores, axis-1, keepdimsTrue)) weights weights / weights.sum(axis-1, keepdimsTrue) output weights V # (n, d) return output def measure_time(n, d64): np.random.seed(42) Q np.random.randn(n, d) K np.random.randn(n, d) V np.random.randn(n, d) # 先跑一次 warmup full_attention(Q, K, V) start time.time() for _ in range(10): full_attention(Q, K, V) avg_time (time.time() - start) / 10 return avg_time for n in [512, 1024, 2048, 4096, 8192]: t measure_time(n) print(f序列长度 {n:6d} - 平均耗时 {t*1000:8.2f} ms)运行结果大致如下具体数值取决于 CPU 性能序列长度 512 - 平均耗时 2.31 ms 序列长度 1024 - 平均耗时 9.05 ms 序列长度 2048 - 平均耗时 36.12 ms 序列长度 4096 - 平均耗时 143.88 ms 序列长度 8192 - 平均耗时 571.20 ms可以看到序列长度翻倍耗时接近翻 4 倍。512 → 1024耗时从 2.31ms 涨到 9.05ms约 3.9 倍4096 → 8192从 143.88ms 涨到 571.20ms约 4 倍。这符合 O(n²) 的预期。上面的代码里scores Q K.T / np.sqrt(d_k)就是那个产生n × n矩阵的关键操作。当n变大时这个矩阵本身就需要 O(n²) 的内存来存储weights V也需要 O(n²) 次乘加运算。真正的成本从这一行就开始了。5. 一个更直观的“档案翻页”实验为了更直观地理解“每生成一个新词都要重翻前面所有记录”可以模拟一个简化版的推理循环。假设我们每次只生成一个新 token但要和前面所有的 token 计算注意力。import numpy as np def generate_one_token(history_Q, history_K, history_V, new_q): 模拟解码器生成一个 token 的过程。 history_Q/K/V: (n, d) 历史 token 的 Q/K/V new_q: (d,) 新 token 的 Query # 新 token 要和历史所有 K 计算注意力分数 scores new_q history_K.T # shape: (n,) weights np.exp(scores - np.max(scores)) weights weights / weights.sum() # 加权求和历史 V得到新 token 的输出 new_output weights history_V # shape: (d,) return new_output def simulate_generation(total_tokens, d64): np.random.seed(0) # 初始历史第一个 token history_K np.random.randn(1, d) history_V np.random.randn(1, d) cumulative_time 0.0 for step in range(total_tokens - 1): new_q np.random.randn(d) start time.time() generate_one_token(history_K, history_V, new_q) cumulative_time time.time() - start # 生成后把新 token 的 K/V 加入历史 new_k np.random.randn(d) new_v np.random.randn(d) history_K np.vstack([history_K, new_k.reshape(1, -1)]) history_V np.vstack([history_V, new_v.reshape(1, -1)]) return cumulative_time import time for total in [100, 500, 1000, 2000]: t simulate_generation(total) print(f生成 {total:5d} 个 token累计注意耗时 {t*1000:8.2f} ms)这个模拟演示了一个关键点每生成一个新 token历史长度都会增加导致下一次注意力计算更慢。累计下来生成n个 token 的总耗时接近 O(n²)。这里还没有包含 KV Cache 用于解决 K/V 重算的优化但仍然体现了“新词要和所有历史记录交互”的核心成本。无论你用什么框架、什么优化器注意力分数矩阵的n × n增长都是绕不开的物理成本。KV Cache 可以避免重复计算 K 和 V但新 token 和旧 token 的交互仍然需要逐对计算。6. 既然全注意力贵那为什么还要用听起来全注意力又慢又贵为什么不直接换掉其实是因为它在表达能力上有不可替代的优势。全注意力让每一个位置都能直接访问整个序列的任意位置信息通路极短路径长度只有 1。相比之下早期的循环神经网络RNN和卷积神经网络CNN在处理长距离依赖时信息需要经过很多步才能从一个位置传递到另一个位置传递过程中容易衰减和丢失。比如句子“小明昨天去了北京他今天早上给我打电话说……”代词“他”要指代“小明”这两个词之间隔了很长的距离。全注意力可以直接建立“他”和“小明”的联系而 RNN 需要一步一步地把信息传过去传着传着可能就丢了。这就是为什么 Transformer 能在大规模语料上取得更好效果的核心原因之一。所以“全注意力贵”不是一个需要被否定的缺点而是一个需要被“管理”的成本。工业界的思路从来不是简单丢弃全注意力而是在保留表达能力的同时用各种方式降低它的计算和存储开销。7. 工业界降本的四条路线理解了全注意力的成本模型后再看工业界的优化方案会清晰很多。目前主流路线可以分成四类。7.1 稀疏注意力只算一部分两两关系稀疏注意力的核心思路是不是每个 token 都需要和其他所有 token 都建立联系我们可以根据位置模式只保留一部分注意力分数。典型代表是 Local Attention滑动窗口注意力每个 token 只和它附近的w个 token 计算注意力复杂度变成 O(n × w)其中w是窗口大小。因为大部分语言现象确实只依赖局部上下文局部窗口往往就能覆盖大部分需求。Longformer、BigBird 都采用了类似的思路。问题也很明显如果某个重要信息出现在窗口之外模型就无法直接建立联系。虽然 BigBird 里加入了一些全局 token 和随机连接来弥补但本质上仍然是一个信息通路受限的系统。7.2 线性注意力用核函数近似线性注意力的核心思路是把 softmax 注意力中的exp(Q K^T)分解为φ(Q) φ(K)^T这样矩阵乘法的结合顺序可以被改变把 O(n²) 的复杂度降成 O(n)。具体来说原始的注意力输出是Output softmax(Q K^T) V其中 softmax 里的exp是逐元素计算的无法和矩阵乘法直接结合。如果用一个函数φ把 Q 和 K 各自映射到新的空间让exp(Q K^T) ≈ φ(Q) φ(K)^T那么Output ≈ φ(Q) (φ(K)^T V)注意这里的运算顺序先算φ(K)^T V这是一个(n, d)和(d, d)的矩阵乘法复杂度 O(n × d²)再算φ(Q)和它的乘积同样是 O(n × d²)。整个注意力计算变成了关于n的线性复杂度。这就是 Linear Attention 这个名称的由来。它把注意力计算从 O(n²) 降到了 O(n)。从标题里看Kimi Linear 的 Linear 很可能就是沿着这条技术路线在做优化用线性复杂度的注意力来替代传统的全注意力从而让长文本处理成本大幅降低。当然线性注意力也有代价核函数的近似不一定能精确模拟 softmax 的注意力分布某些任务上效果会有损失。这也是为什么工业界往往不是简单地替换而是做混合。7.3 Flash Attention不改变算法改变实现Flash Attention 的路线完全不同。它不做任何近似而是通过 IO 感知的 CUDA kernel 设计把注意力计算的中间结果尽量留在 SRAM 里减少和 HBM显存之间的数据搬运从而大幅提升速度、降低显存占用。它有两个关键技巧Tiling分块把 Q、K、V 切成小块在 GPU 的 SRAM 上完成局部计算避免把n × n的完整分数矩阵写回显存Online Softmax在分块计算的同时用一个 online 的方式更新 softmax 的归一化项保证结果和标准 softmax 完全一致。这意味着 Flash Attention 能在结果不变的前提下让显存占用从 O(n²) 降到 O(n)速度也大幅提升。今天的主流大模型训练基本都离不开 Flash Attention。它的限制在于它优化的是实现效率不是算法复杂度。序列极长时O(n²) 的理论复杂度仍然存在最终还是会遇到瓶颈。7.4 综合方案多级混合实际产品通常不是只用一种方案。更常见的做法是底层用 Flash Attention 做标准全注意力的高效实现某些层用稀疏注意力或线性注意力降成本推理阶段用 KV Cache、PagedAttention 等管理显存。这是一个系统工程问题。算法结构、算子实现、内存管理、调度策略每一层都能省一点合起来才有质的改变。8. Kimi Linear从标题能看到什么信号目前关于 Kimi Linear 的公开信息还不多。这里只能基于标题做合理推理不构成对具体产品的定论。从命名习惯看“Linear”大概率指 Linear Attention也就是线性注意力机制。这类方案要解决的核心问题恰恰就是全注意力机制在长文本场景下的二次方成本。如果 Kimi Linear 确实沿着线性注意力的方向优化它可能在解决这么几件事长文本处理成本下降把序列长度对计算量的影响从二次方降到线性长上下文的边际成本大幅降低单次推理的显存占用下降线性注意力不需要显式构建n × n的注意力矩阵KV 缓存压力也可能更小长对话场景的延迟优化每一步生成不再需要和越来越长的历史逐一计算延迟随上下文增长的趋势会更平缓。但要强调的是线性注意力方案在部分任务上可能存在精度损失需要在效果和效率之间做取舍。真正的产品方案一定会在算法层、工程层做大量补偿和调优。对开发者来说关注 Kimi Linear 这类方案的意义在于如果你正在做长文档处理、多轮对话、代码仓库级上下文分析或者需要把大模型部署到显存有限的机器上线性注意力的思路很可能直接关系到你的成本和体验。9. 什么时候需要关注全注意力的成本优化不是所有场景都需要立刻优化全注意力成本。你可以用下面这张表来对照判断场景序列长度是否需要优化短文本分类、情感分析几百 token 以内全注意力完全够用无需额外优化中等长度摘要、翻译1K ~ 4K tokenFlash Attention 级别优化足够长文档问答、多轮对话8K ~ 32K token强烈建议关注注意力成本优化代码仓库理解、超长文档分析32K token 以上必须考虑线性注意力或稀疏注意力判断标准很简单当序列长度超过 4K 以后注意力分数矩阵的规模就开始从“可以接受”变成“明显吃力”。超过 8K 后如果不做任何优化光是注意力计算就会成为系统瓶颈。10. 实践建议如果你要评估一个长文本方案工程师在选型时往往容易被“支持超长上下文”的宣传吸引。这里给你几个可落地的评估思路。10.1 先测延迟曲线再测精度不要只看模型能不能处理 128K 上下文要看它在 128K 上下文下的延迟。真正该做的是固定相同的 prompt 前缀分别测试 1K、4K、16K、32K 下的生成首 token 延迟。如果延迟曲线接近线性说明方案的有效性较高如果延迟在某个长度后突然暴涨说明它的注意力实现可能还是二次方瓶颈。10.2 监控 KV Cache 的变化推理阶段最有价值的监控指标不是显存总量而是 KV Cache 占用的显存变化。如果序列长度增加时KV Cache 增长过快说明方案在存储上并没有做到真正的优化。10.3 用小规模任务验证效果线性注意力之类的高效方案通常会在某些任务上有精度损失。上线前用你自己的业务数据构造长文本测评集对比标准全注意力和优化方案的输出质量。如果损失在可接受范围内效率收益就值得拿如果业务对精度要求极高就尽量选择不改变算法结果的方案比如 Flash Attention。11. 常见问题与排查思路问题现象可能原因排查方式解决方案长文本训练时 OOM注意力分数矩阵占用显存过大用torch.cuda.memory_summary()查看显存分配确认n × n矩阵是否被实例化换用 Flash Attention 或梯度检查点生成速度随上下文增长明显变慢每次生成都要和全部历史计算注意力对比不同上下文长度下的首 token 延迟绘制延迟曲线使用 KV Cache、PagedAttention或评估线性注意力方案KV Cache 显存占用过高缓存了所有层的 K/V计算理论 KV Cache 大小对比实际显存占用减少上下文长度或使用分组查询注意力 GQA换成线性注意力后效果下降核函数近似不够精确在小规模数据集上做 A/B 测试对比生成质量考虑混合方案只在部分层使用线性注意力推理速度上去了但首批 token 延迟很高Prefill 阶段仍然在算全量注意力区分 prefill 和 decode 阶段分别测时对 prefill 阶段做 chunked prefill 或算法级优化12. 工程落地建议与注意事项如果你决定在你的项目里尝试降低注意力成本的方案有几点工程经验值得注意。12.1 先区分训练和推理的瓶颈训练阶段瓶颈更多在计算量本身推理阶段瓶颈经常在显存带宽和 KV Cache 容量。同一个优化手段在训练和推理里的收益可能完全不同。比如 Flash Attention 在训练阶段收益明显推理阶段如果已经用了 KV Cache提升幅度就会小一些。12.2 尽量复用成熟实现不要自己从头实现稀疏注意力或线性注意力。Hugging Face Transformers、FlashAttention、xFormers 等库已经有大量经过验证的实现。自己实现的版本可能在精度、数值稳定性、边界处理上踩很多坑。12.3 做效果回归不要只看速度任何复杂度优化都必须配合效果回归。建议准备一个包含长文本依赖任务的数据集例如“从第 5000 个 token 里找答案”的阅读理解任务。这类任务最能反映模型是否真的学会了长距离建模。12.4 注意上下文长度的实际效果有些模型宣传支持超长上下文但实际使用时可能只对长文档中的局部片段敏感。上线前用你自己的数据做“定位式”测试比跑一个通用榜单更有参考价值。13. 总结与后续学习方向全注意力机制之所以贵不是因为某个具体实现写得差而是它从设计上就要让每个 token 和所有 token 建立两两联系。这种“每个新词都要重翻百万页记录”的做法带来了 O(n²) 的计算复杂度和 O(n) 的 KV Cache 存储成本。在序列长度超长后这两项成本都会让人难以承受。解决这个问题的方向大致有四条稀疏注意力、线性注意力、Flash Attention 这类 IO 优化以及综合性的工程方案。Kimi Linear 大概率属于第二条路线也就是通过线性复杂度的注意力来降低长文本处理成本。至于最终效果如何要看它在算法精度和工程效率之间的平衡做得怎么样。如果你接下来想深入这个方向可以按顺序做三件事用文中的最小代码示例跑一遍把全注意力的复杂度曲线画出来建立直觉阅读 Linear Attention 的经典论文理解核函数近似的数学原理在你的业务场景里选择一个小规模任务对比标准注意力和优化方案的实际差异。技术的取舍永远是效率与质量的平衡。理解全注意力为什么贵并不是为了抛弃它而是为了在正确的场景选择正确的工具。
返回列表