Transformer的架构由编码器(Encoder) 和解码器(Decoder) 两大部分组成
这里重点介绍编码器部分的核心机制:自注意力机制(Self-Attention),它是Transformer的核心创新之一。
TIP工程上,行优先,常见于代码,例如X∈Rn×d;列优先,常见于数学公式。这里为了方便理解,采用列优先的方式。两者在乘积顺序和转置上面有一点区别,但本质上是一样的。
| 符号 | 含义 |
|---|
| n | 输入序列的长度(词元个数) |
| d | 词嵌入向量的维度 |
| X∈Rd×n | 输入序列的词嵌入矩阵,每一列是一个词元的嵌入向量 |
| Q∈Rdk×n | 查询(Query)矩阵,每一列是一个词元的查询向量 |
| qi∈Rdk | 第 i 个词元的查询向量 |
| K∈Rdk×n | 键(Key)矩阵,每一列是一个词元的键向量 |
| kj∈Rdk | 第 j 个词元的键向量 |
| V∈Rdv×n | 值(Value)矩阵,每一列是一个词元的值向量 |
| vj∈Rdv | 第 j 个词元的值向量 |
| dk | 查询和键向量的维度 |
| dv | 值向量的维度 |
| 中间量 | 含义 |
|---|
| sij | 第 i 个词元对第 j 个词元的注意力得分 |
| S∈Rn×n | 注意力得分矩阵 |
| A∈Rn×n | 注意力权重矩阵 |
| Aij 或 Attentionij 或 αij | 第 i 个词元对第 j 个词元的注意力权重 |
| Z∈Rdv×n | 自注意力机制的输出矩阵 |
| zi∈Rdv | 第 i 个词元的输出向量 |
| 参数 | 含义 |
|---|
| WQ∈Rdk×d | 查询的权重矩阵 |
| WK∈Rdk×d | 键的权重矩阵 |
| WV∈Rdv×d | 值的权重矩阵 |
| WO∈Rd×hdv | 多头注意力共享的输出矩阵;混合拼接后的各头输出并映射回 d 维嵌入空间 |
”自” 参注意力#
在机器翻译中,传统的注意力(Cross-Attention,交叉注意力)是:解码器在生成“苹果”这个词时,去编码器里找输入句子“I love apples”中哪个词(I / love / apples)最相关。
而自注意力中的“自”指的是:Query(查询)、Key(键)、Value(值)这三个向量,全部来自同一个输入序列本身。
词嵌入向量/矩阵#
词嵌入向量(Embedding):每个输入词都会被映射为一个高维向量,称为 词嵌入向量。
假设输入序列长度为 n,每个词的嵌入维度为 d,则输入序列可以表示为一个 词嵌入矩阵 X=[x1,x2,...,xn]∈Rd×n。
QKV向量/矩阵#
为了计算注意力,需要为每个词生成 三个 不同的向量。这三个向量是通过三个 可训练 的 权重矩阵 与 输入向量 相乘得到的,对于一整个输入序列也就是三个 权重矩阵 与 词嵌入矩阵 X 相乘:
-
Query(查询)矩阵 Q=WQX=[WQx1,WQx2,...,WQxn]=[q1,q2,...,qn]∈Rdk×n,
- 其中 WQ∈Rdk×d 是查询的权重矩阵,dk 是查询向量的维度。
- “提问者”。这个词想知道:“在当前语境下,我应该关注谁?”
-
Key(键)矩阵 K=WKX=[WKx1,WKx2,...,WKxn]=[k1,k2,...,kn]∈Rdk×n,
- 其中 WK∈Rdk×d 是键的权重矩阵。
- “被问者”。这个词说:“我的特征是XXX,看看你要找的是不是我?”
TIP这里可以发现,Query和Key的维度是一样的,都是 dk,这是因为注意力得分是通过 Query 和 Key 的点积计算的,点积要求两个向量的维度相同。
Q 和 K 都把词嵌入向量投影到一个低维的查询空间(Query Space)和键空间(Key Space),这个低维空间的维度通常比原始嵌入维度小很多(例如 dk=128<d=12888)。
值得注意的是,当键向量 kj 与查询向量 qi “对齐” 的时候,意味着第 i 个token 的嵌入在当前语境下应该更多地 关注 第 j 个token的嵌入。用点积来衡量这种对齐程度是一个非常自然的选择,因为点积在向量空间中可以反映两个向量的相似性。
为了数值稳定性,通常会将点积结果除以 dk,这是因为在高维空间中,向量的点积可能会变得非常大(若 q,k 分量方差近似为 1,则点积方差随 dk 增长;除以 dk 使尺度稳定),从而导致 Softmax 函数的梯度消失问题。
后续的注意力权重矩阵 A,就是对这个矩阵 按行 做 Softmax 归一化。(如果一开始是工程中常见的行优先表示法,那么就是按列做 Softmax 归一化)
- Value(值)矩阵 V=WVX=[WVx1,WVx2,...,WVxn]=[v1,v2,...,vn]∈Rdv×n,
- 其中 WV∈Rdv×d 是值的权重矩阵,dv 是值向量的维度。
- “实际内容”。一旦确认了“提问者”和“被问者”很匹配,这就是我实际要传递给你的具体语义信息。
这里 Wv 的输入输出空间都是 d 维的嵌入空间,对于单头 dv=d,也就是值向量的维度和原始嵌入维度相同。但是这里可以采用 低秩分解 的技巧拆成两个小矩阵降低参数量和计算量
self-attention计算步骤(单头)#
第1步:计算注意力得分(点积)#
计算任意两个词元 i 和 j 的注意力得分(标量)
sij=qiTkj=l=1∑dkqilkjl得分矩阵 S=[sij]∈Rn×n,其中 sij 表示第 i 个词元对第 j 个词元的注意力得分。
S=QTK得到的矩阵图(注意力模式,Attention Pattern):
Query · Key 的相似度矩阵点击上方任一词元,观察它作为 Query 时与所有 Key 的点积
第2步:缩放 + Softmax(按行做 Softmax,使每一行之和为 1)#
A=softmax(dkS)A是注意力权重矩阵,A∈Rn×n,其中 Aij 表示第 i 个词元对第 j 个词元的注意力权重。
第3步:加权求和(对应的Value向量)#
对于每个词元 i,其输出向量 zi 是所有 值向量 以 注意力权重 为系数的加权和:
zi=j=1∑nAijvj
用注意力权重读取 Value选择一个 Query,查看它如何从所有 Value 中汇聚信息
这个向量 zi 是第 i 个 token 从上下文读取到的 单头自注意力输出,维度为 dv。它描述了其他 token 应向当前位置传递什么信息;但它一般还不是可以直接加到 xi∈Rd 上的更新量,因为通常 dv=d。标准多头注意力会先拼接所有头的输出,再通过 WO 映射回 d 维,随后才进行残差相加。
只有在单头且明确取 dv=d、并省略输出投影时,才可以把它简化理解为:
xi←xi+zi写成矩阵形式:
Z=VAT∈Rdv×n矩阵形式#
Z=VAT=V⋅softmax(dkQTK)T
多头注意力机制(Multi-Head Attention)#
多头 指的是:将查询、键、值向量分别映射到 多个子空间 中,进行多次注意力计算,然后将结果拼接起来。
有点类似于CNN中的多通道卷积,每个通道可以学习到不同的特征表示。这里每个头也可以看作是一个独立的注意力机制,它们可以关注输入序列的不同方面。
除了每个头各自拥有的 WQr,WKr,WVr 外,整个多头注意力模块还有一套 共享 参数 WO。它不是第四种 Q/K/V 投影,也不直接作用于输入 X:每个头先完成注意力加权并得到低维输出,所有头的输出拼接后,才由 WO 将它们混合并写回原始嵌入空间。
整体数据流可以概括为:
XWQr,WKr,WVrQr,Kr,VrAttentionZrConcatZconcatWOΔX+XXattn其中 r=1,…,h 表示第 r 个头;前三个投影矩阵决定每个头“如何查询、如何匹配、如何传递内容”,而 WO 决定如何将所有头读出的信息重新组合成对嵌入的更新。
每个头的QKV以及输出#
对于每个头 r,有独立的权重矩阵 WQr,WKr,WVr,
维度:WQr,WKr∈Rdk×d,WVr∈Rdv×d。
对于上下文中的每个位置,也就是每个 token 的嵌入,每个头都会计算自己的低维上下文信息 Zr∈Rdv×n:
- 行数 dv:这个头为每个 token 读出的特征;
- 列数 n:序列中的 token 位置。
Zr=VrArT=Vr⋅softmax(dk(Qr)T(Kr))T拼接(Concatenate)所有头的输出#
将 h 个头的输出矩阵在行方向(特征维度)上堆叠:dv 维的输出向量拼接成 hdv 维的向量,列数仍然是 n:
Zconcat=[Z1;Z2;...;Zh]∈Rhdv×n也就是说:token 的位置(列)不动,只把不同头为这个 token 提供的特征接在一起。
TIP在常见的行优先代码表示中,张量形状是 n×dv,所以会说“沿最后一个维度 拼接”;本质完全相同:固定 token 维度,拼接特征维度。
输出回到原始维度,并通过残差更新#
需要一个新参数矩阵 WO∈Rd×hdv,将拼接后的输出映射回原始嵌入空间。将这个映射结果记为注意力子层提供的更新量 ΔX:
ΔX=WOZconcat∈Rd×n因此,对第 i 个位置,有:
Δxi=WOzi(1)⋮zi(h),xiattn=xi+Δxi这里 WO 会混合各个头读出的信息,并将其翻译回原始的 d 维嵌入空间;残差连接则保留原始表示,只叠加注意力带来的上下文增量。实际 Transformer Block 还会结合归一化层:原始 Transformer 常写为 LayerNorm(X+ΔX)(Post-Norm),现代 LLM 中更常见的是先归一化再计算注意力,即 X+MHA(Norm(X))(Pre-Norm)。
小巧思#
对于嵌入维度 d,通常选择 dk=dv=d/h,这样每个头的输出维度为 dv,拼接后总维度为 hdv=d,与输入维度一致。
- 不增加总参数量。
- 这种“降维投影 + 多头并行”的设计,强迫每个头必须在低维空间(64 维)里寻找特征。由于每个头的初始权重随机且独立训练,它们会自然演化出不同的关注重点(有的擅长局部纹理,有的擅长全局形状)。
工程实现中的参数化变体#
上文使用的是最清晰、也是原始 Transformer 中最常见的写法:每个词嵌入向量 xi∈Rd 分别经过三套独立参数,得到 Query、Key 和 Value:
qi=WQxi,ki=WKxi,vi=WVxi这里 WQ∈Rdk×d、WK∈Rdk×d、WV∈Rdv×d,三者互不共享参数。实际工程中,为了降低参数量、计算量或 KV Cache 的显存占用,常会改变 QKV 的参数化方式;但注意力的基本计算逻辑仍是“用 Q 与 K 计算权重,再用权重加权 V”。
每头 Value 的低秩分解与输出矩阵#
在单头的直观理解中,可以把 Value 看成一个从嵌入空间映射回嵌入空间的完整线性变换。但在标准多头注意力中,每个头不会先产生 d 维 Value 再加权求和,而是先投影到较小的 dv 维空间。对于第 r 个头,记这个下投影为:
WV↓(r)∈Rdv×d,V(r)=WV↓(r)X其注意力加权结果为:
H(r)=V(r)(A(r))T=WV↓(r)X(A(r))T∈Rdv×n若从“每个头都提出一个 d 维嵌入更新”的角度理解,还可以为该头引入一个上投影:
WV↑(r)∈Rd×dv,ΔX(r)=WV↑(r)H(r)于是这个头概念上的完整 Value 映射是:
WV,full(r)=WV↑(r)WV↓(r)∈Rd×d由于中间维度 dv 通常远小于 d,这个完整映射的秩最多为 dv;这就是常说的“Value 的低秩分解”。它并不改变语义:WV↓(r) 负责从词嵌入中抽取该头需要传递的内容,WV↑(r) 负责把这份内容映射为对原始嵌入空间的更新。
实际实现不会为每个头分别执行上投影。先将所有头的低维输出在特征维度拼接:
Hconcat=[H(1);H(2);…;H(h)]∈Rhdv×n再使用一个整个多头模块共享的输出矩阵:
ΔX=WOHconcat,WO∈Rd×hdv将 WO 按列分块,可把它理解为把各头的上投影“钉”在一起:
WO=[WV↑(1)WV↑(2)…WV↑(h)]因此,上式等价于 ΔX=∑r=1hWV↑(r)H(r)。论文和代码中,单个头的 WV(r) 通常就是这里的 WV↓(r);所有“Value-up”合并后对应的是 WO。这也解释了为什么前文的标准写法是先得到 Zconcat,再乘一次 WO,而不是为每个头单独乘一个 d×dv 矩阵。
共享 QKV 的低秩中间投影(可选变体)#
一种可选的低秩参数化是先把词嵌入压缩到一个维度较小的中间表示:
ci=Bxi,B∈Rr×d,r≪d再从这个共享表示分别生成三类向量:
qi=AQci,ki=AKci,vi=AVci其中 AQ∈Rdk×r、AK∈Rdk×r、AV∈Rdv×r。把两步合并后,有:
WQ=AQB,WK=AKB,WV=AVB因此,Q、K、V 共享的是右侧的低维基底 B,即它们都先从 xi 中抽取同一个压缩特征 ci,再由不同的 AQ,AK,AV 赋予“查询 / 匹配 / 传递内容”三种语义。
WQ,WK,WV 共同采用了一个共享的低秩分解。这样能节省参数,但也会限制三种投影可独立表达的信息,因此是否采用取决于模型规模和效果需求。
共享 KV:MQA 与 GQA#
另一个常见工程优化并不是低秩分解,而是减少多头注意力中 Key 和 Value 的副本数。设有 h 个 Query 头:
qi(r)=WQ(r)xi,r=1,…,h
- MQA(Multi-Query Attention):所有 Query 头共用同一组 K,V,即 ki=WKxi、vi=WVxi。这样解码时只需缓存一份 K/V。
- GQA(Grouped-Query Attention):将 h 个 Query 头分成若干组;同一组共享一组 K,V。它在 MHA(每头独立 KV)与 MQA(全部头共享 KV)之间折中。
两者中,Query 仍然由自己的 WQ(r) 产生;被共享的是 K/V。主要收益是降低自回归推理时的 KV Cache 显存和带宽开销。
潜变量 KV 压缩#
还有一类做法会把每个词的 Key/Value 先压缩为潜变量,例如:
ciKV=WDKVxi,ki=WUKciKV,vi=WUVciKV其中 ciKV∈RrKV 是用于缓存的低维 KV 潜变量,WDKV 是下投影矩阵,WUK 与 WUV 分别恢复 Key 和 Value。这类设计的重点是缓存 ciKV 而非完整的 ki,vi,从而降低 KV Cache。Query 可以保持独立投影,也可以使用另一套低秩分解。