Transformer 的计算复杂度:为什么序列变长后注意力会变重
从 QKᵀ 的注意力矩阵出发,理解 Transformer 中注意力和前馈网络的主要计算成本,并区分训练、生成与 KV 缓存带来的不同开销。
相关工具
为什么要单独看计算复杂度
模型参数量告诉我们需要保存多少可学习数值,但它不能完整说明一次输入要花多少计算。Transformer 的成本还和序列长度、模型维度、前馈维度以及生成方式有关。尤其在长文本场景中,序列长度会让注意力部分迅速变重。
理解复杂度不是为了背一堆公式,而是为了回答几个实际问题:为什么输入加长后速度变慢,为什么显存会突然不够,为什么训练和生成的成本表现不同,以及为什么长上下文需要专门的结构优化。
最重要的主线是:自注意力要建立位置之间的两两关系,前馈网络则对每个位置分别计算。两种计算的增长方式不同。

Q 与 K 的矩阵乘法产生 N×N 的注意力分数矩阵,序列长度增加时,两两关系数量按平方增长。
QKᵀ 为什么会产生 N×N 矩阵
假设输入序列有 N 个位置,每个位置经过投影后得到 Q 和 K。Q 的形状可以写成 N×d_k,K 的形状也是 N×d_k。计算 QKᵀ 时,K 被转置成 d_k×N,于是结果形状是 N×N。
结果矩阵中的第 i 行、第 j 列表示第 i 个 Query 位置对第 j 个 Key 位置的相关性分数。因为每个位置都要和所有位置比较,所以一共需要考虑 N² 个位置对。
这就是标准全注意力随序列长度平方增长的根源。序列从 1,000 个词元增加到 2,000 个词元时,位置对数量不是增加一倍,而是接近增加到原来的四倍。
注意力不只有 QKᵀ 一次矩阵乘法
得到注意力分数后,还要进行缩放、掩码和 Softmax,再用归一化后的权重与 V 相乘。权重矩阵的形状仍是 N×N,V 的形状是 N×d_v,最终得到 N×d_v 的输出。
因此,注意力的主要成本通常来自两类矩阵乘法:QKᵀ 生成位置两两关系,以及注意力权重乘 V 汇总上下文。它们都要接触 N×N 的中间关系,序列越长,计算和内存压力越明显。
多头注意力会把维度分到多个头,但在总模型维度相近时,分头主要改变表示方式,并不会消除全序列两两交互的基本成本。
前馈网络为什么通常按序列长度线性增长
前馈网络对每个位置独立使用同一套线性变换。输入有 N 个位置,就相当于把同一个两层网络应用 N 次。每个位置内部的计算量与 d_model 和 d_ff 有关,但不同位置之间不需要构造 N×N 的关系矩阵。
因此,固定模型维度时,前馈网络的主要计算量通常随 N 线性增加。序列变长会让它变重,但不会像全注意力那样产生平方增长。
在某些模型配置下,短序列上的主要成本可能来自前馈网络;当序列足够长时,注意力的两两关系会逐渐成为更突出的瓶颈。具体谁占主导,要结合 d_model、d_ff 和 N 一起判断。

注意力需要位置之间全对全交互,主要项随序列长度平方增长;前馈网络逐位置计算,主要项随序列长度线性增长。
内存成本为什么同样重要
计算量大意味着需要更多运算,但显存压力还来自中间激活和注意力矩阵。训练时通常要保存反向传播所需的中间结果,N×N 的注意力分数或权重矩阵可能占据大量空间。
这也是长序列训练中容易出现显存不足的原因之一。即使参数量没有变化,输入长度增加也会让激活值和注意力相关张量变大。减小批次、使用梯度检查点、分块计算或更省内存的注意力实现,都可以缓解部分压力。
不要把“能不能装下模型参数”和“能不能处理这段长度的输入”混为一谈。前者是模型权重和优化器状态的问题,后者还要考虑激活与注意力中间结果。
训练阶段可以并行处理整段序列
训练时,目标序列通常已经完整存在,编码器或解码器可以把整段输入一次性送进矩阵运算。解码器虽然使用因果掩码,但掩码只限制可见范围,不要求硬件按时间步等待,所以多个位置仍然可以并行计算。
全序列并行提高了训练吞吐,但它会一次性产生整段序列的激活和注意力中间结果。序列越长,单次计算的并行规模越大,同时内存压力也越高。
训练效率常常要在序列长度、批次大小和显存之间平衡。把多个短序列拼成批次可以利用并行硬件,但填充位置过多又会浪费计算,因此数据整理也会影响实际效率。
生成阶段为什么会逐步变慢
自回归生成每次只新增一个词元,但新位置需要关注前面已经生成的内容。随着上下文变长,当前 Query 仍然要和越来越多的 Key 比较,因此单步计算会随已生成长度增加。整段生成的累计成本也会不断叠加。
如果每一步都把完整前缀重新送入所有层,系统会反复计算历史位置的 Key 和 Value,产生大量重复工作。实际推理通常会缓存历史 Key 和 Value,让新位置只计算自己的投影,并与缓存中的历史信息交互。
KV 缓存不能让长序列完全没有成本。缓存本身会占用显存,而且每生成一个新词仍要读取越来越长的历史 Key 和 Value。它主要减少重复计算,不能消除上下文长度带来的读取和注意力成本。

训练可以全序列并行,生成按词元逐步推进,并用 KV 缓存复用历史 Key 和 Value,减少重复计算。
KV 缓存到底缓存什么
在解码器的自注意力中,历史位置已经算过的 Key 和 Value 可以保留下来。下一轮生成新位置时,只需要为新位置计算 Query、Key、Value,把新的 Key 和 Value 追加到缓存,再用当前 Query 查询完整的历史 Key。
缓存通常按层、按头保存,因此上下文越长、层数越多、模型维度越宽,KV 缓存占用的内存越大。批量生成多个序列时,每条序列的缓存也需要分别维护。
KV 缓存解决的是重复计算问题,不是把所有历史内容压缩成一个固定大小的向量。历史 Key 和 Value 仍然会随着序列增长,长上下文推理的内存管理因此非常重要。
长上下文为什么需要额外设计
直接把最大序列长度不断调大,会同时增加注意力计算、激活内存和 KV 缓存压力。位置编码也要能够在更长位置上保持有区分度,训练数据还需要包含足够多的长文本,模型才有机会学会如何利用远处信息。
常见的优化思路包括稀疏注意力、滑动窗口、分块处理、局部与全局注意力组合,以及对 KV 缓存进行压缩或分页管理。它们的共同目标是减少不必要的全序列两两交互,或者降低历史信息的保存成本。
长上下文不只是把窗口变大。模型是否真正能利用远处信息,还取决于训练方式、位置表示、注意力模式和任务数据。可处理的长度与有效理解的长度并不完全相同。
如何判断计算瓶颈在哪里
先看序列长度。如果 N 很大,优先关注注意力矩阵和激活内存;如果 N 较短但 d_model 或 d_ff 很大,前馈网络和线性投影可能占据主要计算。再看是训练还是生成,两者的并行方式和缓存策略不同。
其次区分计算瓶颈和内存瓶颈。有些场景不是算得慢,而是频繁读取参数或 KV 缓存;有些场景则是注意力矩阵计算本身占用了大量算力。只看总参数量,无法定位问题。
最后结合实际指标观察:单步生成延迟、每秒词元数、显存占用、批次吞吐和输入长度变化后的曲线。复杂度公式提供方向,真实硬件和实现方式决定最终表现。
确认 N 是否让注意力矩阵成为主要成本。
检查 d_model 与 d_ff 对线性变换的影响。
区分训练全序列并行和生成逐步推进。
把参数、激活、注意力矩阵和 KV 缓存分开观察。
把复杂度放回 Transformer 主线
Transformer 的核心优势是用注意力直接建立远距离关系,但这种全局交互也带来了 N×N 的计算和内存成本。前馈网络提供逐位置的非线性加工,计算随序列长度线性增长;两者共同构成一层的主要计算来源。
训练时,矩阵运算让整段序列可以并行处理;生成时,自回归依赖让系统必须逐步推进,KV 缓存则负责尽量减少历史重复计算。理解这三条线,就能解释很多速度和显存现象。
以后看到长上下文、稀疏注意力、Flash Attention 或 KV 缓存等概念,都可以先问它们在解决哪一种成本:减少两两交互、降低中间内存,还是减少生成阶段的重复计算。
常见问题
为什么自注意力的复杂度常写成 O(N²)?
因为每个位置通常要和序列中的所有位置计算相关性,QKᵀ 会产生 N×N 的注意力分数矩阵。
KV 缓存能把生成复杂度变成常数吗?
不能。KV 缓存减少了历史 Key 和 Value 的重复计算,但新 Query 仍需要读取不断增长的历史信息,缓存本身也会占用内存。
序列变长时,前馈网络也会平方增长吗?
标准逐位置前馈网络对每个位置独立计算,固定模型维度时主要随序列长度线性增长;平方增长主要来自全注意力的位置两两交互。