🎯 课程主题
交叉注意力机制(Cross Attention)——解码器与编码器信息融合的桥梁。
📝 核心知识点
1. 交叉注意力的定义
- 概念说明:Q 来自解码器,K/V 来自编码器的注意力机制。
- 关键细节:
- 与自注意力机制本质相同,区别仅在输入来源
- 自注意力:Q、K、V 均来自同一输入
- 交叉注意力:Q 来自解码器,K、V 来自编码器
- 也可构成多头注意力机制
2. 交叉注意力的输入来源
- 概念说明:交叉注意力融合编码器和解码器的信息。
- 关键细节:
- K、V 来源:编码器的输出(如 "我喜欢打篮球" 的特征)
- Q 来源:解码器前一子层的输出(因果掩码注意力的输出)
- 编码器信息经 K、V 进一步提取
- 与解码器信息做残差融合
3. 交叉注意力的计算流程
- 概念说明:Q 与 K 计算注意力得分,再与 V 融合。
- 关键细节(以中译英为例):
- 编码器输入:$A_1, A_2, A_3, A_4$("我喜欢打篮球")
- 解码器输入:$X_1, X_2, X_3, X_4$(
<S> I like playing) - $A_i \times W_K = K_i$(K 来自编码器)
- $A_i \times W_V = V_i$(V 来自编码器)
- $X_j \times W_Q = Q_j$(Q 来自解码器)
- $Q_j \times K_i$ → 注意力得分 $\alpha_{ji}$
- Softmax 归一化 → 注意力分配
- 与 $V$ 相乘 → 信息融合输出 $B_j$
4. 交叉注意力的意义
- 概念说明:交叉注意力使解码器能关注编码器的相关信息。
- 关键细节:
- 输出既包含编码器提取的信息(源语言)
- 也包含解码器的信息(目标语言已生成部分)
- 解码器的 Q 帮助从编码器信息中提取相关内容
- 实际信息仍来自编码器的 V
5. 三种注意力机制总结
- 概念说明:Transformer 中有三种注意力机制。
- 关键细节:
- 自注意力(编码器):Q、K、V 均来自编码器,可带填充掩码
- 因果掩码注意力(解码器):Q、K、V 均来自解码器,带因果掩码
- 交叉注意力(解码器):Q 来自解码器,K、V 来自编码器,可带填充掩码
- 三者本质相同,区别在输入来源和掩码类型
🧮 核心公式与推导
交叉注意力公式:
$$\text{CrossAttention}(Q_{dec}, K_{enc}, V_{enc}) = \text{Softmax}\left(\frac{Q_{dec} K_{enc}^T}{\sqrt{d_k}}\right) V_{enc}$$
其中:
- $Q_{dec} = X_{dec} W_Q$:Q 来自解码器
- $K_{enc} = X_{enc} W_K$:K 来自编码器
- $V_{enc} = X_{enc} W_V$:V 来自编码器
注意力得分:
$$\alpha_{ji} = \text{Softmax}\left(\frac{Q_j \cdot K_i}{\sqrt{d_k}}\right)$$
表示解码器第 $j$ 个位置对编码器第 $i$ 个位置的关注程度。
信息融合:
$$B_j = \sum_i \alpha_{ji} V_i$$
物理意义:解码器的每个位置通过 Q 查询编码器的信息,按关联程度加权提取编码器的特征。
🏗️ 模型架构与数据流向
💻 代码实战
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
class CrossAttention(nn.Module):
"""交叉注意力机制"""
def __init__(self, d_model, d_k):
super().__init__()
self.W_Q = nn.Linear(d_model, d_k, bias=False) # 用于解码器输入
self.W_K = nn.Linear(d_model, d_k, bias=False) # 用于编码器输入
self.W_V = nn.Linear(d_model, d_k, bias=False) # 用于编码器输入
self.d_k = d_k
def forward(self, dec_input, enc_output, src_mask=None):
"""
dec_input: [batch, tgt_len, d_model] 解码器输入 (Q来源)
enc_output: [batch, src_len, d_model] 编码器输出 (K,V来源)
src_mask: [batch, src_len] 编码器padding mask
"""
# Q来自解码器, K/V来自编码器
Q = self.W_Q(dec_input) # [batch, tgt_len, d_k]
K = self.W_K(enc_output) # [batch, src_len, d_k]
V = self.W_V(enc_output) # [batch, src_len, d_k]
# 注意力得分: [batch, tgt_len, src_len]
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(self.d_k)
# 填充mask (屏蔽编码器中的padding位置)
if src_mask is not None:
scores = scores.masked_fill(src_mask.unsqueeze(1) == 0, float('-inf'))
attn_weights = F.softmax(scores, dim=-1)
output = torch.matmul(attn_weights, V) # [batch, tgt_len, d_k]
return output, attn_weights
# 示例: 中译英
d_model = 512
d_k = 64
batch_size = 2
cross_attn = CrossAttention(d_model, d_k)
# 编码器输出 (源语言: 我喜欢打篮球, 5个token)
enc_output = torch.randn(batch_size, 5, d_model)
# 解码器输入 (目标语言: <S> I like playing, 4个token)
dec_input = torch.randn(batch_size, 4, d_model)
# 编码器padding mask (假设第2个样本最后1位是padding)
src_mask = torch.tensor([
[1, 1, 1, 1, 1], # 样本1: 全部有效
[1, 1, 1, 1, 0] # 样本2: 最后1位是padding
])
output, attn = cross_attn(dec_input, enc_output, src_mask)
print(output.shape) # torch.Size([2, 4, 64])
print(attn.shape) # torch.Size([2, 4, 5]) 解码器4个位置对编码器5个位置的关注度
⚠️ 常见问题与避坑指南
- Q 来自解码器,K/V 来自编码器,不要搞反
- 交叉注意力得分矩阵是 $L_{tgt} \times L_{src}$(非方阵)
- 交叉注意力也可带填充掩码(屏蔽编码器的 padding)
- 交叉注意力不带因果掩码(解码器可看到编码器全部信息)
💡 个人总结与延伸
交叉注意力是连接编码器和解码器的桥梁,使解码器能根据当前生成状态查询编码器的源语言信息。这是 Encoder-Decoder 架构(如原始 Transformer、T5、BART)的核心。现代 Decoder-only 架构(如 GPT 系列)不使用交叉注意力,但理解它对掌握完整 Transformer 架构至关重要。