深度学习的数学原理(二十九)—— 不同注意力的讨论
在前几篇文章中,我们已经系统走完了注意力机制的完整脉络:从最基础的缩放点积注意力,到解决多语义问题的多头注意力,再到防止"偷看未来"的掩码注意力,最后到Transformer核心的编码器-解码器交叉注意力。
我们已经形成了两个最基本的认知:
- 自注意力:Q、K、V全部来自同一个序列,用于建模序列内部的依赖关系
- 标准交叉注意力:Q来自解码器,K、V来自编码器,用于生成时对齐源序列信息
但顺着这个逻辑自然会产生一个非常本质的问题:
如果Q、K、V每一个都可以独立选择来自编码器或者解码器,理论上会有2×2×2=8种组合。
这8种组合是否都能正常运行?是否都有实际应用价值?
特别是"Q来自解码器、KV来自编码器"和"Q来自编码器、KV来自解码器"这两种完全相反的结构,为什么作用天差地别?
本文将从矩阵乘法的底层约束出发,先筛选出所有数学上合法的组合,再逐一分析其语义、作用与差异,最后通过数值算例和代码验证所有结论。
一、数学约束:矩阵乘法的维度要求
在讨论任何组合之前,我们必须先回到缩放点积注意力的原始数学定义,这是所有讨论的基础:
Attention(Q,K,V)=softmax(QK⊤dk)V \text{Attention}(Q,K,V) = \text{softmax}\left(\frac{QK^\top}{\sqrt{d_k}}\right)V Attention(Q,K,V)=softmax(dkQK⊤)V
整个计算分为两个不可分割的步骤,每一步都有严格的矩阵维度约束:
第一步:计算相似度分数 QK⊤QK^\topQK⊤
- 设 Q∈RLQ×dkQ \in \mathbb{R}^{L_Q \times d_k}Q∈RLQ×dk(查询序列,长度LQL_QLQ,特征维度dkd_kdk)
- 设 K∈RLK×dkK \in \mathbb{R}^{L_K \times d_k}K∈RLK×dk(键序列,长度LKL_KLK,特征维度dkd_kdk)
- 转置后 K⊤∈Rdk×LKK^\top \in \mathbb{R}^{d_k \times L_K}K⊤∈Rdk×LK
- 矩阵乘法要求:Q的列数 = K的列数(特征维度相同)
- 运算结果:分数矩阵 S∈RLQ×LKS \in \mathbb{R}^{L_Q \times L_K}S∈RLQ×LK
第二步:加权求和 W⋅VW \cdot VW⋅V
- 经过softmax后,权重矩阵 W=softmax(S)∈RLQ×LKW = \text{softmax}(S) \in \mathbb{R}^{L_Q \times L_K}W=softmax(S)∈RLQ×LK
- 设 V∈RLV×dvV \in \mathbb{R}^{L_V \times d_v}V∈RLV×dv(值序列,长度LVL_VLV,特征维度dvd_vdv)
- 矩阵乘法强制要求:权重矩阵的列数 = V的行数
- 也就是必须满足:LK=LV\boldsymbol{L_K = L_V}LK=LV
核心结论(专栏定理)
K和V必须来自同一个序列,长度完全一致,位置一一对应。
这是注意力机制能够正常运行的必要且充分条件,没有任何变通空间。
只有Q可以独立选择来源。
因此,所谓的"8种自由组合"中,K和V来源不同的4种组合是完全非法的,根本无法完成矩阵运算,更谈不上应用价值。
(值得一提,Transformers本就是为了解决变长序列,因此很难保证解码器和编码器的序列永远等长,因此除非非常极限的情况,来自不同序列的KV才有可能可以相乘)
二、所有合法的注意力组合:共4种
基于上述约束,注意力的合法结构只由两个独立变量决定:
- Q的来源:编码器序列XXX 或 解码器序列YYY
- KV的来源:编码器序列XXX 或 解码器序列YYY
总共有且仅有 4种合法组合,全部满足矩阵乘法约束,全部有清晰的物理意义和实际应用价值。
我们用统一的二元组(Q来源,KV来源)(Q\text{来源}, KV\text{来源})(Q来源,KV来源)表示:
| 序号 | 组合 | 标准名称 | 核心语义 | 典型应用 |
|---|---|---|---|---|
| 1 | (X, X) | 编码器自注意力 | 源序列自己查询自己,建模源文本上下文 | Transformer编码器 |
| 2 | (X, Y) | 反向交叉注意力 | 源序列查询目标序列,吸收目标信息 | 双向翻译、多模态对齐 |
| 3 | (Y, X) | 标准交叉注意力 | 目标序列查询源序列,抽取源信息 | 机器翻译、文本生成 |
| 4 | (Y, Y) | 解码器自注意力 | 目标序列自己查询自己,建模时序依赖 | Transformer解码器(带掩码) |
三、核心问题:对称组合为什么作用截然不同?
最容易让人困惑的就是这一对互为镜像的结构:
- 标准交叉注意力:(Y,X)(Y, X)(Y,X)
- 反向交叉注意力:(X,Y)(X, Y)(X,Y)
它们只是交换了Q和KV的来源,看上去像是"正反两面",但实际功能却完全不同。下面我们从直观理解和数学证明两个层面彻底讲透。
3.1 直观理解:谁向谁提问,信息往哪里流
标准交叉注意力 (Y, X)
- Q = Y:正在生成的目标词发出查询(“我现在要生成英文,该看中文的哪个词?”)
- KV = X:在源序列中匹配,并从源序列提取信息
- 信息流向:源序列X → 目标序列Y
- 作用:生成目标语言时,精准对齐源文本,把源语言的语义搬运到目标语言,这是机器翻译的核心。
反向交叉注意力 (X, Y)
- Q = X:源序列的词发出查询(“我这个中文词,对应已经生成的英文的哪个位置?”)
- KV = Y:在目标序列中匹配,并从目标序列提取信息
- 信息流向:目标序列Y → 源序列X
- 作用:编码源文本时,对齐已生成的目标文本,把目标语言的信息反向融入源语言,多用于双向翻译、端到端语音识别。
3.2 数学证明:二者不可能等价
我们给出三条严格的数学证明,从不同维度说明它们的本质差异。
证明1:权重矩阵的维度与空间不可逆
-
标准交叉注意力的权重矩阵:
WY→X=softmax(YX⊤dk)∈RLY×LX W_{Y\rightarrow X} = \text{softmax}\left(\frac{YX^\top}{\sqrt{d_k}}\right) \in \mathbb{R}^{L_Y \times L_X} WY→X=softmax(dkYX⊤)∈RLY×LX
每一行对应一个目标位置,每一列对应一个源位置,表示"第i个目标词对第j个源词的关注程度"。 -
反向交叉注意力的权重矩阵:
WX→Y=softmax(XY⊤dk)∈RLX×LY W_{X\rightarrow Y} = \text{softmax}\left(\frac{XY^\top}{\sqrt{d_k}}\right) \in \mathbb{R}^{L_X \times L_Y} WX→Y=softmax(dkXY⊤)∈RLX×LY
每一行对应一个源位置,每一列对应一个目标位置,表示"第i个源词对第j个目标词的关注程度"。
两个矩阵形状互逆、定义域和值域完全不同,不存在任何线性变换可以将一个转化为另一个,数学上就是两个完全不同的算子。
证明2:输出语义由KV唯一决定
注意力的本质可以拆分为两个独立的部分:
-
QK部分:决定在哪里分配权重(对齐空间)
-
V部分:决定加权什么内容(信息空间)
-
标准交叉注意力的V来自X → 输出结果完全携带源序列的语义信息
-
反向交叉注意力的V来自Y → 输出结果完全携带目标序列的语义信息
即便QK计算出的相似度分数完全相同,只要KV来源不同,输出的语义信息就会彻底不同。
证明3:梯度传播方向相反,优化目标不同
在训练过程中,梯度的流向由Q和KV的来源决定:
-
标准交叉注意力:损失从解码器输出 → 交叉注意力层 → 梯度回传给编码器的KV
→ 优化目标:让编码器的输出更适合被解码器查询,更容易被对齐。 -
反向交叉注意力:损失从编码器输出 → 反向交叉注意力层 → 梯度回传给解码器的KV
→ 优化目标:让解码器的输出更适合被编码器查询,更容易被对齐。
优化目标完全不同,因此学到的参数分布和模型功能必然不同。
四、具体数值算例:手算验证差异
我们沿用本专栏一贯的极简风格,用"我爱→I love"这个迷你翻译任务来手算验证,所有数值都可以直接口算。
设定:
- 编码器源序列X:[“我”, “爱”],对应的特征向量:
X=[1.02.03.04.0]X = \begin{bmatrix}1.0 & 2.0 \\ 3.0 & 4.0\end{bmatrix}X=[1.03.02.04.0](LX=2,dk=2L_X=2, d_k=2LX=2,dk=2) - 解码器目标序列Y:[“I”, “love”],对应的特征向量:
Y=[0.51.52.53.5]Y = \begin{bmatrix}0.5 & 1.5 \\ 2.5 & 3.5\end{bmatrix}Y=[0.52.51.53.5](LY=2,dk=2L_Y=2, d_k=2LY=2,dk=2) - 缩放因子:dk=2≈1.414\sqrt{d_k} = \sqrt{2} \approx 1.414dk=2≈1.414
例1:标准交叉注意力 (Y, X)
计算生成"I"时对源序列的注意力:
Q1=Y[0]=[0.5,1.5]Q1X⊤=[0.5,1.5]⋅[1.03.02.04.0]=[3.5,7.5]缩放分数=[3.5/1.414,7.5/1.414]≈[2.475,5.303]W=softmax([2.475,5.303])≈[0.056,0.944]O1=0.056×[1.0,2.0]+0.944×[3.0,4.0]≈[2.888,3.888] \begin{align*} Q_1 &= Y[0] = [0.5, 1.5] \\ Q_1X^\top &= [0.5, 1.5] \cdot \begin{bmatrix}1.0 & 3.0 \\ 2.0 & 4.0\end{bmatrix} = [3.5, 7.5] \\ \text{缩放分数} &= [3.5/1.414, 7.5/1.414] \approx [2.475, 5.303] \\ W &= \text{softmax}([2.475, 5.303]) \approx [0.056, 0.944] \\ O_1 &= 0.056 \times [1.0, 2.0] + 0.944 \times [3.0, 4.0] \approx [2.888, 3.888] \end{align*} Q1Q1X⊤缩放分数WO1=Y[0]=[0.5,1.5]=[0.5,1.5]⋅[1.02.03.04.0]=[3.5,7.5]=[3.5/1.414,7.5/1.414]≈[2.475,5.303]=softmax([2.475,5.303])≈[0.056,0.944]=0.056×[1.0,2.0]+0.944×[3.0,4.0]≈[2.888,3.888]
可以看到,生成"I"时主要关注源序列的第二个词"爱",输出携带源序列的信息。
例2:反向交叉注意力 (X, Y)
计算编码"我"时对目标序列的注意力:
Q1=X[0]=[1.0,2.0]Q1Y⊤=[1.0,2.0]⋅[0.52.51.53.5]=[3.5,9.5]缩放分数=[3.5/1.414,9.5/1.414]≈[2.475,6.718]W=softmax([2.475,6.718])≈[0.014,0.986]O1=0.014×[0.5,1.5]+0.986×[2.5,3.5]≈[2.472,3.472] \begin{align*} Q_1 &= X[0] = [1.0, 2.0] \\ Q_1Y^\top &= [1.0, 2.0] \cdot \begin{bmatrix}0.5 & 2.5 \\ 1.5 & 3.5\end{bmatrix} = [3.5, 9.5] \\ \text{缩放分数} &= [3.5/1.414, 9.5/1.414] \approx [2.475, 6.718] \\ W &= \text{softmax}([2.475, 6.718]) \approx [0.014, 0.986] \\ O_1 &= 0.014 \times [0.5, 1.5] + 0.986 \times [2.5, 3.5] \approx [2.472, 3.472] \end{align*} Q1Q1Y⊤缩放分数WO1=X[0]=[1.0,2.0]=[1.0,2.0]⋅[0.51.52.53.5]=[3.5,9.5]=[3.5/1.414,9.5/1.414]≈[2.475,6.718]=softmax([2.475,6.718])≈[0.014,0.986]=0.014×[0.5,1.5]+0.986×[2.5,3.5]≈[2.472,3.472]
可以看到,编码"我"时主要关注目标序列的第二个词"love",输出携带目标序列的信息。
五、PyTorch代码验证
下面给出严格符合数学约束的通用注意力实现,验证上述算例的结果。
import torch
import torch.nn.functional as F
def standard_attention(Q, KV, d_k):
"""
严格合规的通用注意力实现
约束:K和V来自同一个序列,输入为统一的KV张量,位置一一对应
输入:Q[batch, L_Q, d_k], KV[batch, L_K, d_k]
输出:加权结果,注意力权重
"""
K = KV
V = KV
attn_scores = torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(d_k, dtype=torch.float32))
attn_weights = F.softmax(attn_scores, dim=-1)
output = torch.matmul(attn_weights, V)
return output, attn_weights
# 构造与手算相同的输入
d_k = 2
X = torch.tensor([[[1.0, 2.0], [3.0, 4.0]]], dtype=torch.float32) # 源序列:我、爱
Y = torch.tensor([[[0.5, 1.5], [2.5, 3.5]]], dtype=torch.float32) # 目标序列:I、love
# 1. 标准交叉注意力 (Q=Y, KV=X)
out_cross, w_cross = standard_attention(Y, X, d_k)
# 2. 反向交叉注意力 (Q=X, KV=Y)
out_rev, w_rev = standard_attention(X, Y, d_k)
# 打印结果
print("="*60)
print("标准交叉注意力 (生成I时对源序列的注意力)")
print(f"注意力权重:{w_cross[0,0].numpy()}")
print(f"输出结果:{out_cross[0,0].numpy()}")
print("="*60)
print("反向交叉注意力 (编码我时对目标序列的注意力)")
print(f"注意力权重:{w_rev[0,0].numpy()}")
print(f"输出结果:{out_rev[0,0].numpy()}")
print("="*60)
运行结果

六、总结
本文从矩阵乘法的底层约束出发,完整推导了注意力QKV来源的所有合法组合,回答了开头提出的所有问题:
- 8种组合中只有4种合法:K和V必须来自同一个序列,否则无法完成矩阵乘法,其余4种组合在数学上和工程上都不成立。
- 对称组合作用不同的根本原因:Q决定对齐方向,KV决定信息内容,二者来源互换会导致对齐空间、信息流向、梯度传播路径全部反转,数学上不可逆、功能上不等价。
- 应用价值:经典Transformer采用了(X,X)、(Y,Y)、(Y,X)三种结构,保证了结构简洁与翻译任务的有效性;而(X,Y)这类反向结构在双向建模、多模态融合、语音文本对齐等场景中具有不可替代的价值。
更多推荐



所有评论(0)