Transformer 架构梳理

4,677 字 · 13 min

Transformer 是 2017 年《Attention Is All You Need》提出的序列建模架构,现在的大语言模型基本都建立在它上面。这篇按从零件到整机的顺序过一遍:文本怎么变成向量,attention 怎么算,多头、位置编码、一个 block 的组成,encoder-decoder 到 decoder-only 的演变,训练目标,最后是推理阶段的采样、KV cache 和性能瓶颈。

读完应该能回答这几个问题:

  • attention 为什么能替代 RNN,付出了什么代价
  • 一个 Transformer block 里有几个矩阵乘法,参数主要在哪,怎么从超参数估算总参数量
  • 现在的大模型为什么都是 decoder-only
  • KV cache 存的是什么,为什么它是长上下文的显存大头,推理为什么卡在带宽上

为什么需要它

Transformer 出现之前,序列建模的主力是 RNN 和它的变体 LSTM、GRU。它们最大的问题是串行:第 t 步的隐状态依赖第 t−1 步,长度为 n 的序列必须一步一步算 n 次,GPU 再多也用不上。

另一个问题是长距离依赖。信息从第 1 个 token 传到第 100 个 token,中间要经过 99 次状态更新,每次都会被压缩和覆盖一部分。LSTM 的门控缓解了这个问题,但没有解决。

Transformer 的思路是不再按时间步逐个传递,而是让序列里任意两个位置直接建立联系。这样一来,任意两个 token 之间的路径长度都是 O(1),整个序列的所有位置也能同时计算,前面两个问题就一起解决了。

每层计算量 串行步数 任意两位置的路径长度
RNN O(n · d²) n O(n)
Self-attention O(n² · d) 1 O(1)

代价也在这张表里:任意两个位置之间都要算一次关系,计算量和 n 的平方成正比,序列长度翻一倍,attention 的开销就翻四倍。后面长上下文的种种麻烦,大多是从这里来的。

从文本到向量

模型不直接处理字符。文本先经过 tokenizer 切成 token,主流做法是 BPE 或它的变体:从单个字符出发,反复把语料里最常相邻出现的两个片段合并成一个新 token,直到词表达到预定大小。常见词会成为一个 token,生僻词被切成几个子词,任何字节序列都能被表示出来,不存在词表外的情况。词表大小从 GPT-2 的 5 万到现在常见的 10 到 15 万。

每个 token 有一个整数 id,通过一张 V × d_model 的 embedding 矩阵查表,变成一个 d_model 维的向量。V 是词表大小,d_model 是模型的隐藏维度,7B 量级的模型一般是 4096。模型末端还有一个 d_model × V 的输出层,把最后的向量映射回词表上的分布。这两个矩阵形状互为转置,很多模型直接共用一份权重。

从这一步开始,模型内部流动的就是长度为 n 的一串 d_model 维向量。位置信息还没有加进来,后面单独讲。

Self-attention

attention 从每个 token 的向量 x 出发做三个线性变换:

1
q = x · W_Q      k = x · W_K      v = x · W_V

W_Q、W_K、W_V 是三个 d_model × d_k 的矩阵,整个序列共用同一组,位置 1 和位置 1000 用的是同一个 W_Q。这一点既是它能并行的前提,也是它没有位置概念的原因。

三个向量可以这样理解:q 是「我在找什么」,k 是「我有什么」,v 是「我要传出去的内容」。可以把整个过程想成一次「软」的查表。普通哈希表用 key 精确匹配,取出一个 value;attention 用 q 和所有 k 算相似度,再按相似度把所有 v 加权平均,没有命中和未命中之分,只有权重大小。

具体地说,一个 token 拿自己的 q 和序列里所有 token 的 k 做点积,得到一组相关性分数,分数过 softmax 变成和为 1 的权重,再用这组权重对所有 v 做加权求和,结果就是这个 token 的新表示。把整个序列的 q、k、v 各堆成矩阵,一次矩阵乘法就能把所有 token 一起算完:

1
Attention(Q, K, V) = softmax( Q Kᵀ / √d_k ) · V

scaled dot-product attention 的计算流程

有几处细节值得多说两句。

Q Kᵀ 是一个 n × n 的矩阵,第 i 行第 j 列是第 i 个 token 对第 j 个 token 的关注分数。softmax 按行做,第 i 行归一化之后,就是第 i 个 token 分给序列里每个位置的权重。n² 的计算量和显存占用都来自这个矩阵。

除以 √d_k 的原因要从数值范围说起。假设 q 和 k 的各分量独立、零均值、单位方差,那么点积 q·k 是 d_k 个乘积之和,方差就是 d_k。d_k 取 64 或 128 时,点积的数值会很大,而 softmax 对大数值很敏感,最大的那个分量会拿走几乎全部权重,其余趋近于 0,梯度也跟着消失。除以 √d_k 把方差拉回 1,softmax 才能工作在有梯度的区间。原论文对这一步只有一句话,但省掉它的实现经常训练不收敛。

mask 是可选的一步,在 softmax 之前把某些位置的分数置为 −∞,softmax 之后这些位置的权重就精确为 0。后面 decoder 用的 causal mask 和 padding 的处理都靠它。

输出的形状和输入一样,n 个 token 进去,n 个 d_v 维向量出来,每个位置的输出都是全序列 v 的加权平均。所以 attention 可以一层层叠加,每一层在上一层的基础上重新聚合。

最小实现十几行,只用 numpy,对着公式看一遍就够:

1
2
3
4
5
6
7
8
9
10
11
import numpy as np

def attention(Q, K, V, mask=None):
d_k = Q.shape[-1]
scores = Q @ K.T / np.sqrt(d_k) # n × n
if mask is not None:
scores = np.where(mask, scores, -np.inf)
scores -= scores.max(axis=-1, keepdims=True) # 数值稳定
w = np.exp(scores)
w /= w.sum(axis=-1, keepdims=True) # softmax,按行
return w @ V # n × d_v

Multi-head

单头 attention 有一个局限:softmax 出来的只有一种权重分布,一个 token 在一层里只能用一种方式聚合信息。但语言里同时存在很多种关系,语法上的主谓、指代上的先行词、位置上的邻近,一种权重分布照顾不过来。

多头的做法是把 d_model 切成 h 份,每份 d_model/h 维,各自独立做一次 attention,把 h 个结果拼接起来,再过一个线性层 W_O 融合:

1
2
head_i = Attention(x W_Q^i, x W_K^i, x W_V^i)
MultiHead(x) = Concat(head_1, …, head_h) · W_O

每个头的维度 d_head = d_model / h,现在的模型基本固定在 128:7B 模型 4096 维配 32 个头,70B 模型 8192 维配 64 个头。实现上 h 个头并不是分别算的,而是把 h 组 W_Q 拼成一个 d_model × d_model 的大矩阵,一次乘法得到所有头的 q,再 reshape 成 h × n × d_head,在头这个维度上批量做 attention。

参数量和单头一样:h 个头的投影矩阵各是 d_model × d_head,加起来还是 d_model × d_model。多头换来的是同一层里 h 种不同的关注模式。事后可视化可以看到,有的头在追句法结构,有的头专门看前一个 token,有的头在做指代消解,也有相当一部分头看起来没干什么,剪掉之后效果几乎不变。

位置编码

attention 有一个容易忽略的性质:它是置换等变的。把输入序列打乱顺序,输出也跟着打乱,但每个 token 得到的表示完全不变。也就是说 attention 本身没有位置的概念,「我爱你」和「你爱我」在它看来一样。这是「整个序列共用同一组 W」的直接后果。

所以位置信息必须显式地加进去,做法大致经历了三代。

第一代是原论文的正弦编码。给每个位置 pos 生成一个 d_model 维向量,偶数维用 sin、奇数维用 cos,频率随维度指数递减:

1
2
PE(pos, 2i)   = sin( pos / 10000^(2i/d) )
PE(pos, 2i+1) = cos( pos / 10000^(2i/d) )

直接加在输入 embedding 上。低维度变化快,高维度变化慢,和二进制计数各位的规律类似。它不需要学,任意长度都能算,但模型要自己从这堆三角函数里学会「相对位置」,学得并不好。

第二代是可学习的绝对位置编码,BERT 和 GPT-2 用的就是它,给每个位置一个可训练的向量,简单直接。问题是训练时最长见过 1024,推理时第 1025 个位置就没有向量了,完全没法外推。

第三代是 RoPE,旋转位置编码,LLaMA、Qwen 等现在的主流模型都用它。思路换了一下,不再往输入上加东西,而是在算 attention 分数之前对 q 和 k 做旋转。把 d 维向量两两分组看成 d/2 个二维向量,第 i 组用一个固定的角频率 θ_i,位置 m 的 token 把这一组旋转 m·θ_i:

1
2
[x1']   [ cos(mθ)  −sin(mθ) ] [x1]
[x2'] = [ sin(mθ) cos(mθ) ] [x2] θ_i = 10000^(−2i/d)

二维旋转有个很好的性质:两个向量分别旋转 α 和 β 之后再做点积,结果只和 α − β 有关。于是位置 m 的 q 和位置 n 的 k 做点积,天然只依赖 m − n,相对位置信息直接进了 attention 分数,不需要模型自己去学。低频组编码远距离关系,高频组编码近距离关系,分工和正弦编码一样,只是塞进去的位置从输入换到了 q 和 k。

RoPE 的外推能力比绝对编码好,但也有限。位置超出训练长度后,高频维度的旋转角度会进入训练时没见过的区间。现在常见的长上下文扩展方法,比如 NTK-aware 缩放、YaRN,做的都是调整那组基频 θ_i,让长位置对应的角度落回模型熟悉的范围。另一条路是 ALiBi,什么编码都不做,直接在 attention 分数上按距离减一个线性惩罚,外推表现好,表达力弱一些,用得少。

一个 block 长什么样

把上面的零件装起来就是一个 Transformer block。以现在通用的 Pre-LN 结构为例:

Pre-LN 的 Transformer block

1
2
x  = x + MHA( Norm(x) )
x' = x + FFN( Norm(x) )

两条支路,每条都是先归一化,过子层,再加回残差。

FFN 是两层全连接加一个非线性,先升维到 4·d_model,激活,再降回 d_model。它对每个位置独立计算,位置之间不交换信息,交换信息是 attention 的事。现在的主流模型把它换成了 SwiGLU,多一条门控支路,中间维度相应调到 8/3·d_model 左右,参数量持平。有一种流行的解读把 FFN 看成 key-value 存储:第一层的每一行是一个 key,和输入做内积得到激活强度,第二层的对应列是 value,按强度加权取出。按这个解读,attention 负责在 token 之间搬运信息,FFN 负责存储和变换,模型记住的事实主要在 FFN 里。

参数分布可以直接算。attention 的四个投影矩阵各是 d²,合计 4·d²;FFN 两层各 4·d²,合计 8·d²。一个 block 大约 12·d²,三分之二在 FFN。整个模型的参数量也就有了估算公式:

1
总参数 ≈ 12 · d² · L  +  V · d       (L 层,V 是词表大小)

代入 LLaMA-7B 的 d = 4096、L = 32、V = 32000:12 × 4096² × 32 约 64 亿,加上 embedding 的 1.3 亿,约 66 亿,实际是 67 亿,差的部分来自 SwiGLU 多出来的那条支路。这个公式反过来也有用,看到一个模型的参数量,大致就能推出它的 d 和 L。

残差给梯度留了一条直通路径,几十上百层的网络能训练起来全靠它。它也提供了一个看模型的角度:残差流是一条贯穿所有层的 d_model 维主干,每一层从里面读一点、往里面写一点修正,没有哪一层会彻底重写表示。

归一化方面,原论文用 LayerNorm,对每个位置的 d 维向量减均值、除标准差,再乘一个可学习的缩放 g、加一个偏移 b。现在基本都换成了 RMSNorm,去掉减均值和偏移项,效果几乎一样,算得更快:

1
2
LayerNorm(x) = (x − μ) / σ · g + b
RMSNorm(x) = x / √( mean(x²) + ε ) · g

Pre-LN 和 Post-LN 的区别在归一化的位置。原论文是 Post-LN,先加残差再归一化,这种结构在深层时训练不稳定,必须配合 learning rate warmup。Pre-LN 把归一化挪到子层之前,残差路径上没有 Norm,梯度更平稳,可以用更大的学习率,现在是标配。

把 L 个这样的 block 叠起来,前面接 embedding,后面接一次归一化和输出层,就是完整的模型。GPT-3 是 96 层,LLaMA-70B 是 80 层。

从 Encoder-Decoder 到 Decoder-only

原论文是为翻译设计的,结构分两半。encoder 双向看完整个源句,每个位置可以关注前后所有位置;decoder 生成目标句,用 causal mask 保证只能看到已经生成的部分,另外多一个 cross-attention 去看 encoder 的输出,q 来自 decoder,k 和 v 来自 encoder。

后来这两半各自成了流派。BERT 只留 encoder,双向,适合分类、抽取这类理解任务。T5 保留完整的 encoder-decoder,至今在翻译、摘要上还有使用。GPT 只留 decoder,去掉 cross-attention,只做一件事:给定前文,预测下一个 token。

decoder-only 最终胜出,主要有三个原因。训练目标极其简单,任何文本天然就是训练数据,不需要标注。理解和生成统一到了同一个框架里,回答问题就是续写问题后面的文字。规模化也最顺,scaling law 的曲线在 decoder-only 上最干净。

causal mask 是这个结构的核心。把 Q Kᵀ 矩阵的上三角全部置为 −∞,softmax 后权重为 0,第 i 个位置就只能看到 1 到 i。它还带来一个训练上的好处:一个长度为 n 的序列,一次前向传播同时得到 n 个位置的预测,每个位置都在预测它的下一个 token,相当于 n 个训练样本并行完成。这种训练方式叫 teacher forcing,喂给模型的永远是真实的前文,而不是它自己生成的。

causal mask 与 KV cache

训练目标

预训练只有一个损失函数:每个位置对下一个 token 的预测,和真实的下一个 token 之间的交叉熵。输出层给出词表上的 logits,softmax 之后取真实 token 对应的概率,取负对数,n 个位置求平均。没有别的监督信号。模型的全部能力,包括事实、推理、代码,都是从「把下一个 token 猜准」这一件事里来的。

预训练的数据量级是万亿 token。之后的指令微调(SFT)和偏好对齐(RLHF、DPO)用的是完全相同的架构和损失形式,只是数据从网页换成了对话样本,规模小几个数量级。对齐没有在架构里加任何新东西,它只是在预训练权重上继续做梯度下降。

推理

训练时序列是并行的,推理时不是。生成是自回归的:给一段前文,算出下一个 token 的分布,从中取一个,拼到前文后面,再算下一个。每一步都是一次完整的前向传播。

取哪一个由采样策略决定。logits 除以 temperature 之后再做 softmax,temperature 越低分布越尖,趋近于永远取最大的那个,越高越平,随机性越强。top-p 只在累计概率达到 p 的那批候选里采样,截掉长尾。temperature 设为 0 就是贪心解码,同样的输入永远得到同样的输出。

最朴素的做法是每一步把整个序列重新算一遍。生成 n 个 token,第 t 步的 attention 是 O(t²),累积起来 O(n³)。这里面有大量重复:在 causal mask 之下,前面 token 的 k 和 v 不依赖任何后面的 token,算过一次之后就不会再变。

KV cache 就是把它们存下来。每一步只为新 token 算 q、k、v,把新的 k、v 追加进缓存,然后用这一个 q 对缓存里所有的 k 做 attention。每步的 attention 从 O(t²) 降到 O(t)。上面那张图里的实线行就是这一步真正在算的部分。

代价是显存。缓存的大小是:

1
2 × 层数 × 序列长度 × d_model × 每个数的字节数        (再乘 batch)

拿 LLaMA-7B 算一下:32 层,d_model 4096,fp16。每个 token 是 2 × 32 × 4096 × 2 B = 512 KB。4k 上下文 2 GB,32k 上下文 16 GB,已经和模型权重本身一个量级。MQA 和 GQA 就是为此出现的:让多个 query 头共享同一组 k、v 头。LLaMA-2-70B 用 8 组 kv 头对应 64 个 query 头,缓存缩小 8 倍,效果只有很小的损失。

推理的瓶颈在显存带宽而不是算力。生成每一个 token,都要把全部模型权重从显存里完整读一遍,但每个权重只参与一次乘加。以 A100 为例,显存带宽 2 TB/s,读一遍 14 GB 的 7B 模型权重要 7 ms,而这些权重对应的浮点运算在 A100 上不到 0.1 ms 就能算完,算力大部分时间在等数据。所以推理优化主要在减少读显存的次数和字节数:把多个请求 batch 在一起分摊权重读取,量化到 int8 或 int4 缩小权重体积,用 GQA 缩小 KV cache。FlashAttention 解决的是另一头的问题,它不把 n × n 的分数矩阵完整写进显存,而是分块在片上缓存里算完 softmax 再往下走,结果精确不变,省掉了最大的一块中间读写。

推理分两个阶段,性质不一样。prefill 处理输入的 prompt,所有 token 并行,和训练一样受算力限制,决定首字延迟。decode 一个一个吐 token,受带宽限制,决定每秒多少字。推理服务通常还会把多个请求共同的前缀,比如同一个 system prompt,对应的 KV cache 缓存起来跨请求复用,叫 prefix caching,省掉重复的 prefill。

小结

回到开头的几个问题:

  • attention 用 O(n²) 的两两连接换掉了 RNN 的串行,路径长度 O(1),整个序列并行
  • 一个 block 有 attention 的 4 个投影和 FFN 的 2 到 3 个矩阵,约 12·d² 个参数,三分之二在 FFN;总量约 12·d²·L + V·d
  • decoder-only 赢在目标简单、数据免标注、理解生成统一、规模化最顺
  • KV cache 存的是每一层每个 token 的 k 和 v,随上下文线性增长;decode 阶段每个 token 都要读一遍全部权重,瓶颈在带宽

架构本身到这里就讲完了。从安全的角度再看一遍这套结构,会发现现在 LLM 的很多攻击面在架构层面就已经定下来了,这部分单独写了一篇:从架构看 LLM 的攻击面。