🎯 课程主题
Transformer 模型的训练过程——利用因果掩码实现并行训练与标签屏蔽。
📝 核心知识点
1. 训练数据与标签
- 概念说明:训练需要输入数据和对应的正确标签。
- 关键细节(以 "我是一条狗" → "I am a dog" 为例):
- 输入数据:我是一条狗
- 标签:I am a dog
- 模型输出与标签计算 Loss → 反向传播更新参数
2. 推理 vs 训练的关键区别
- 概念说明:训练时输入的是正确标签,而非预测结果。
- 关键细节:
- 推理:基于前一步的预测结果作为下一步输入(串行)
- 训练:输入正确标签(Teacher Forcing),并行计算所有位置
- 如果训练时用预测结果,错误会累积,导致模型学习困难
3. 因果掩码(Causal Mask)的作用
- 概念说明:因果掩码防止模型在训练时"偷看"未来信息。
- 关键细节:
- 训练时将所有正确标签一次性输入解码器
- 问题:模型可能直接看到后面的答案
- 解决:因果掩码多头注意力机制
- 因果掩码确保:
- 计算位置 1 的 Loss 时,看不到位置 2、3、4
- 计算位置 2 的 Loss 时,看不到位置 3、4
- 以此类推
4. 并行训练
- 概念说明:因果掩码使训练可以并行计算所有位置的 Loss。
- 关键细节:
- 一次性输入所有标签 → 并行计算
- 同时得到所有位置的 Loss:
- "I" 的 Loss
- "am" 的 Loss
- "a" 的 Loss
- "dog" 的 Loss
- 结束符的 Loss
- 因果掩码保证每个位置只看到当前及之前的信息
5. 反向传播更新参数
- 概念说明:通过 Loss 反向传播更新模型所有参数。
- 关键细节:
- 更新解码器参数:因果掩码注意力、交叉注意力、FFN、线性层
- 更新编码器参数:多头注意力、FFN
- 参数包括权重 $W$ 和偏置 $b$
🧮 核心公式与推导
因果掩码矩阵(下三角矩阵):
$$\text{Mask}_{ij} = \begin{cases} 0 & \text{如果 } i \geq j \text{ (可以看到当前位置及之前)} \ -\infty & \text{如果 } i < j \text{ (不能看到未来位置)} \end{cases}$$
带因果掩码的注意力:
$$\text{Attention}(Q, K, V) = \text{Softmax}\left(\frac{QK^T}{\sqrt{d_k}} + \text{Causal Mask}\right) V$$
训练 Loss 计算:
$$\text{Loss} = -\sum_{t=1}^{T} \log P(y_t | y_{<t}, x)$$
其中 $y_t$ 是第 $t$ 个目标词,$y_{<t}$ 是第 $t$ 个位置之前的词(通过因果掩码保证)。
物理意义:因果掩码确保位置 $t$ 只能关注位置 $\leq t$ 的信息,防止训练时数据泄露。
🏗️ 模型架构与数据流向
因果掩码示意(以 4 个 token 为例):
位置: 1 2 3 4
位置1: ✓ ✗ ✗ ✗ (只能看到自己)
位置2: ✓ ✓ ✗ ✗ (看到1,2)
位置3: ✓ ✓ ✓ ✗ (看到1,2,3)
位置4: ✓ ✓ ✓ ✓ (看到1,2,3,4)
💻 代码实战
import torch
import torch.nn.functional as F
def make_causal_mask(seq_len):
"""
生成因果掩码(下三角矩阵)
seq_len: 序列长度
return: [seq_len, seq_len] mask, 0表示可见, -inf表示不可见
"""
mask = torch.triu(torch.ones(seq_len, seq_len), diagonal=1) # 上三角为1
mask = mask.masked_fill(mask == 1, float('-inf')) # 上三角填-inf
return mask
def causal_masked_attention(Q, K, V, causal_mask=None):
"""
带因果掩码的注意力计算
Q: [batch, L, d_k]
K: [batch, L, d_k]
V: [batch, L, d_v]
"""
d_k = Q.size(-1)
scores = torch.matmul(Q, K.transpose(-2, -1)) / math.sqrt(d_k)
if causal_mask is not None:
scores = scores + causal_mask # 上三角变为-inf
attn_weights = F.softmax(scores, dim=-1)
output = torch.matmul(attn_weights, V)
return output
# 训练过程示例
def train_step(model, src, tgt, optimizer, criterion):
"""
model: Transformer模型
src: [batch, src_len] 源语言
tgt: [batch, tgt_len] 目标语言(含起始符)
"""
# 1. 编码器
memory = model.encode(src)
# 2. 构建因果掩码
tgt_len = tgt.size(1)
causal_mask = make_causal_mask(tgt_len).to(tgt.device)
# 3. 解码器 (输入tgt[:-1], 预测tgt[1:])
# 解码器输入: <S> I am a dog
# 解码器目标: I am a dog <E>
tgt_input = tgt[:, :-1] # 去掉最后一个token
tgt_output = tgt[:, 1:] # 去掉第一个token(起始符)
# 4. 前向传播
logits = model.decode(memory, tgt_input, causal_mask)
# logits: [batch, tgt_len-1, vocab_size]
# 5. 计算Loss (并行计算所有位置)
loss = criterion(
logits.reshape(-1, logits.size(-1)), # [batch*(tgt_len-1), vocab_size]
tgt_output.reshape(-1) # [batch*(tgt_len-1)]
)
# 6. 反向传播
optimizer.zero_grad()
loss.backward()
optimizer.step()
return loss.item()
# 因果掩码示例
mask = make_causal_mask(4)
print(mask)
# tensor([[0., -inf, -inf, -inf],
# [0., 0., -inf, -inf],
# [0., 0., 0., -inf],
# [0., 0., 0., 0.]])
⚠️ 常见问题与避坑指南
- 训练用正确标签,推理用预测结果——这是 Teacher Forcing 策略
- 因果掩码方向:上三角为 $-\infty$(屏蔽未来),下三角为 0(保留过去和当前)
- 训练时所有位置的 Loss 并行计算,效率远高于推理的串行生成
- 如果不加因果掩码,模型会"偷看"答案,训练无效
- 解码器输入与目标错位:输入
<S> I am a dog,目标I am a dog <E>
💡 个人总结与延伸
Transformer 的训练过程通过因果掩码实现了并行训练,同时保证了自回归特性。Teacher Forcing 策略(用正确标签作为输入)加速了训练收敛,但也带来了 Exposure Bias 问题(训练时见到的都是正确输入,推理时见到的是可能有误的预测)。现代大模型通过 Scheduled Sampling、RLHF 等技术缓解这一问题。