Teacher Forcing - 教师强制
1 课程概览
本课讲解 Teacher Forcing(教师强制)训练技巧。在 Seq2Seq 架构中,解码器每次使用真实结果作为输入的一部分,而不是上一个时间步的输出结果。作用是防止"一步错,步步错",类似于小孩做题时第三步错了,后续步骤大概率全错。
2 核心概念与定义
- Teacher Forcing:教师强制,解码器使用真实结果作为输入。
- 真实结果(Ground Truth):正确的标签值。
- 预测结果:模型生成的输出。
- 一步错步步错:一个时间步预测错误,后续时间步大概率全错。
3 模型与算法详解
正常流程(不带 Teacher Forcing)
一个词一个词往外蹦,用上次生成的词作为本次输入。
输入 → 预测词1(正确)→ 预测词2(正确)→ 预测词3(错误)→ 预测词4(大概率错)→ ...
问题:如果词3预测错误,用错误的词3作为输入预测词4,词4大概率错误。
Teacher Forcing 流程
解码器每次使用真实结果作为输入的一部分。
输入 → 预测词1 → 用真实词1作为输入 → 预测词2 → 用真实词2作为输入 → 预测词3 → 用真实词3作为输入 → ...
优点:即使词3预测错误,用真实的词3作为输入预测词4,避免错误传播。
示例说明
标签(真实值):[119, 297, 456, 25, ...]
正常流程:
- 预测词1 = 119(正确)
- 用 119 预测词2 = 297(正确)
- 用 297 预测词3 = 300(错误)
- 用 300 预测词4 = ?(大概率错)
Teacher Forcing:
- 预测词1 = 119(正确)
- 用真实值 119 预测词2 = 297(正确)
- 用真实值 297 预测词3 = 300(错误)
- 用真实值 456 预测词4 = ?(不受词3错误影响)
类比
类似于小孩做数学题,共 8 步,第三步错了,后续步骤大概率全错。
- 正常流程:第三步错了,后续全错
- Teacher Forcing:第三步错了,但用正确答案继续,后续不受影响
4 数学原理与推导
正常流程
$$\hat{y}t = f(y{t-1}, h_{t-1})$$
其中 $\hat{y}_{t-1}$ 是上一时刻的预测值。
Teacher Forcing
$$\hat{y}t = f(y{t-1}, h_{t-1})$$
其中 $y_{t-1}$ 是上一时刻的真实值(Ground Truth)。
区别
| 流程 | 输入 | 说明 |
|---|---|---|
| 正常 | $\hat{y}_{t-1}$(预测值) | 用预测结果作为输入 |
| Teacher Forcing | $y_{t-1}$(真实值) | 用真实结果作为输入 |
5 代码示例
import torch
import torch.nn as nn
def train_with_teacher_forcing(encoder, decoder, input_tensor, target_tensor,
encoder_optimizer, decoder_optimizer, criterion,
max_length=10):
"""
使用 Teacher Forcing 的训练函数
:param input_tensor: 输入张量(英语句子)
:param target_tensor: 目标张量(法语句子,真实值)
"""
encoder_optimizer.zero_grad()
decoder_optimizer.zero_grad()
input_length = input_tensor.size(0)
target_length = target_tensor.size(0)
loss = 0
# 编码
encoder_hidden = encoder.initHidden()
for ei in range(input_length):
encoder_output, encoder_hidden = encoder(input_tensor[ei], encoder_hidden)
# 解码
decoder_input = torch.tensor([[SOS]]) # 开始标志
decoder_hidden = encoder_hidden # 用编码器的隐藏状态初始化
# Teacher Forcing:使用真实结果作为输入
for di in range(target_length):
decoder_output, decoder_hidden, _ = decoder(
decoder_input, decoder_hidden, encoder_outputs
)
# 计算损失
loss += criterion(decoder_output, target_tensor[di])
# Teacher Forcing:用真实结果作为下一步输入
decoder_input = target_tensor[di] # 真实值
# 反向传播
loss.backward()
encoder_optimizer.step()
decoder_optimizer.step()
return loss.item() / target_length
# 对比:不带 Teacher Forcing
def train_without_teacher_forcing(encoder, decoder, input_tensor, target_tensor,
encoder_optimizer, decoder_optimizer, criterion,
max_length=10):
"""不带 Teacher Forcing 的训练函数"""
encoder_optimizer.zero_grad()
decoder_optimizer.zero_grad()
loss = 0
# 编码
encoder_hidden = encoder.initHidden()
for ei in range(input_length):
encoder_output, encoder_hidden = encoder(input_tensor[ei], encoder_hidden)
# 解码
decoder_input = torch.tensor([[SOS]])
decoder_hidden = encoder_hidden
for di in range(target_length):
decoder_output, decoder_hidden, _ = decoder(
decoder_input, decoder_hidden, encoder_outputs
)
loss += criterion(decoder_output, target_tensor[di])
# 不带 Teacher Forcing:用预测结果作为下一步输入
topv, topi = decoder_output.topk(1)
decoder_input = topi.squeeze().detach() # 预测值
loss.backward()
encoder_optimizer.step()
decoder_optimizer.step()
return loss.item() / target_length
6 重难点与易错提醒
- ❗重点:Teacher Forcing 使用真实结果作为输入,不是预测结果。
- ❗重点:作用是防止"一步错,步步错"。
- ⚠️易错:训练时用 Teacher Forcing,预测时不能用(没有真实值)。
- 💡深入理解:Teacher Forcing 加快训练速度,提高稳定性。
- 💡深入理解:预测时只能用上一个时间步的输出作为输入。
7 课堂问答精选
Q1:什么是 Teacher Forcing?
A:教师强制,在 Seq2Seq 架构中,解码器每次使用真实结果作为输入的一部分,而不是上一个时间步的输出结果。
Q2:Teacher Forcing 有什么作用?
A:防止"一步错,步步错"。类似于小孩做数学题,第三步错了,后续步骤大概率全错。Teacher Forcing 用真实结果作为输入,避免错误传播。
Q3:Teacher Forcing 和正常流程有什么区别?
A:正常流程用预测结果作为输入,Teacher Forcing 用真实结果作为输入。正常流程一个时间步错误会影响后续,Teacher Forcing 不受影响。
Q4:训练和预测时如何使用 Teacher Forcing?
A:训练时用 Teacher Forcing(有真实值);预测时不能用 Teacher Forcing(没有真实值),只能用上一个时间步的输出作为输入。
8 本课小结
- Teacher Forcing:解码器使用真实结果作为输入。
- 作用:防止"一步错,步步错"。
- 正常流程:用预测结果作为输入。
- Teacher Forcing:用真实结果作为输入。
- 训练时用,预测时不能用。
9 延伸思考
- Teacher Forcing 有什么缺点?
- 如何在训练和预测之间平滑过渡?