🎯 课程主题
因果掩码注意力机制在训练过程中的应用——通过倒三角掩码实现并行训练与未来信息屏蔽。
📝 核心知识点
1. 因果掩码的核心思想
- 概念说明:训练时输入正确标签,通过因果掩码防止模型看到未来信息。
- 关键细节:
- 训练时将所有正确标签一次性输入解码器
- 因果掩码确保每个位置只能看到当前及之前的位置
- 通过倒三角矩阵(上三角为 $-\infty$)实现
2. 因果掩码的构建
- 概念说明:在注意力得分矩阵的上三角填充 $-\infty$。
- 关键细节:
- 输入标签 $X$ → 映射为 $Q, K, V$
- $Q \times K^T$ 得到注意力得分矩阵
- 上三角(未来位置)填充 $-\infty$
- 经 Softmax:$-\infty \to 0$,未来位置注意力得分为 0
3. 因果掩码的效果
- 概念说明:每个位置只能关注当前位置及之前的信息。
- 关键细节(以 5 个 token
<S> I like playing ball为例):- 位置 1(
<S>):只能看到<S>,看不到I like playing ball - 位置 2(
I):能看到<S> I,看不到like playing ball - 位置 3(
like):能看到<S> I like,看不到playing ball - 位置 5(
ball):能看到所有前面的词
- 位置 1(
4. 信息融合过程
- 概念说明:注意力得分矩阵与 V 相乘,每个位置只融合可见信息。
- 关键细节:
- 位置 1 输出:只包含
<S>的信息 - 位置 2 输出:包含
<S>和I的信息融合 - 位置 5 输出:包含所有 token 的信息融合
- 每个位置的融合都包含注意力得分分配
- 位置 1 输出:只包含
5. 因果掩码的价值
- 概念说明:因果掩码是 Transformer 中最巧妙也最难的算法。
- 关键细节:
- 使训练可以并行计算所有位置的 Loss
- 同时保证自回归特性(不看未来)
- 训练和推理的掩码机制本质相同
🧮 核心公式与推导
因果掩码矩阵(以 5 个 token 为例):
$$\text{Mask} = \begin{bmatrix} 0 & -\infty & -\infty & -\infty & -\infty \ 0 & 0 & -\infty & -\infty & -\infty \ 0 & 0 & 0 & -\infty & -\infty \ 0 & 0 & 0 & 0 & -\infty \ 0 & 0 & 0 & 0 & 0 \end{bmatrix}$$
带因果掩码的注意力:
$$\text{Attention}(Q, K, V) = \text{Softmax}\left(\frac{QK^T}{\sqrt{d_k}} + \text{Causal Mask}\right) V$$
Softmax 后的注意力矩阵:
$$A = \begin{bmatrix} 1 & 0 & 0 & 0 & 0 \ \alpha'{21} & \alpha'{22} & 0 & 0 & 0 \ \alpha'{31} & \alpha'{32} & \alpha'{33} & 0 & 0 \ \alpha'{41} & \alpha'{42} & \alpha'{43} & \alpha'{44} & 0 \ \alpha'{51} & \alpha'{52} & \alpha'{53} & \alpha'{54} & \alpha'{55} \end{bmatrix}$$
其中上三角全为 0,每行和为 1。
物理意义:位置 $i$ 只能关注位置 $\leq i$ 的信息,防止训练时数据泄露。
🏗️ 模型架构与数据流向
因果掩码示意:
位置: 1 2 3 4 5
位置1(S): ✓ ✗ ✗ ✗ ✗
位置2(I): ✓ ✓ ✗ ✗ ✗
位置3(like): ✓ ✓ ✓ ✗ ✗
位置4(play): ✓ ✓ ✓ ✓ ✗
位置5(ball): ✓ ✓ ✓ ✓ ✓
💻 代码实战
import torch
import torch.nn.functional as F
import math
def causal_masked_attention(Q, K, V):
"""
因果掩码注意力计算 (训练时使用)
Q: [batch, L, d_k]
K: [batch, L, d_k]
V: [batch, L, d_v]
"""
batch_size, L, d_k = Q.size()
# 1. 计算注意力得分
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k)
# scores: [batch, L, L]
# 2. 构建因果掩码 (上三角为-inf, 下三角为0)
causal_mask = torch.triu(
torch.ones(L, L, device=Q.device), diagonal=1
).bool()
scores = scores.masked_fill(causal_mask, float('-inf'))
# 3. Softmax归一化 (上三角变为0)
attn_weights = F.softmax(scores, dim=-1)
# 4. 信息融合
output = torch.matmul(attn_weights, V)
return output, attn_weights
# 示例: 训练时输入正确标签 "<S> I like playing ball"
L = 5
d_k = 64
Q = torch.randn(2, L, d_k)
K = torch.randn(2, L, d_k)
V = torch.randn(2, L, d_k)
output, attn = causal_masked_attention(Q, K, V)
print(output.shape) # torch.Size([2, 5, 64])
# 验证: 注意力矩阵上三角为0
print(attn[0])
# tensor([[0.33, 0.00, 0.00, 0.00, 0.00], # 位置1只看自己
# [0.20, 0.25, 0.00, 0.00, 0.00], # 位置2看1,2
# [0.15, 0.18, 0.22, 0.00, 0.00], # 位置3看1,2,3
# [0.10, 0.15, 0.20, 0.18, 0.00], # 位置4看1,2,3,4
# [0.08, 0.12, 0.15, 0.20, 0.25]]) # 位置5看全部
⚠️ 常见问题与避坑指南
- 因果掩码方向:上三角为 $-\infty$(屏蔽未来),下三角为 0(保留过去和当前)
- 训练时输入的是正确标签,不是预测结果(Teacher Forcing)
- 因果掩码使训练可以并行计算所有位置的 Loss
- 因果掩码是 Transformer 中最巧妙也最难的算法,务必理解透彻
💡 个人总结与延伸
因果掩码是 Transformer 训练并行化的关键,通过倒三角掩码确保每个位置只能看到历史信息,从而可以一次性输入所有标签并行训练。这一机制是所有自回归模型(GPT 系列)训练的基础。现代大模型训练中,因果掩码的实现被高度优化(如 Flash Attention 中的掩码融合)以提升效率。