🎯 课程主题
Transformer 模型的推理过程——解码器如何自回归地逐步生成翻译结果。
📝 核心知识点
1. 推理过程概述
- 概念说明:推理是模型训练完成后,将输入翻译为目标语言的过程。
- 关键细节:
- 输入:中文句子 "我是一条狗"
- 输出:英文句子 "I am a dog"
- 解码器基于前一时刻的输出,逐 token 生成
2. 编码器处理
- 概念说明:编码器对输入句子进行特征提取。
- 关键细节:
- 输入 "我是一条狗" → 词嵌入 + 位置编码
- 经多头注意力机制 → 层归一化 + 残差连接
- 经前馈神经网络 → 层归一化 + 残差连接
- 输出编码器信息(包含序列信息与注意力分配信息)
- 编码器信息传入解码器
3. 解码器自回归生成
- 概念说明:解码器基于前一步输出逐步生成翻译。
- 关键细节(以 "我是一条狗" → "I am a dog" 为例):
- 第 1 步:输入
<S>(起始符)+ 编码器信息 → 输出概率分布 → 最大概率词 "I" - 第 2 步:输入
<S> I+ 编码器信息 → 输出 "am" - 第 3 步:输入
<S> I am+ 编码器信息 → 输出 "a" - 第 4 步:输入
<S> I am a+ 编码器信息 → 输出 "dog" - 第 5 步:输入
<S> I am a dog+ 编码器信息 → 输出<E>(结束符)
- 遇到结束符
<E>时,推理结束
- 第 1 步:输入
4. 解码器内部流程
- 概念说明:解码器每步经过多个子层处理。
- 关键细节:
- 因果掩码多头注意力(Masked Multi-Head Attention)
- 交叉注意力(Cross Attention):融合编码器信息
- 前馈神经网络 + 残差连接 + 层归一化
- 线性层 + Softmax → 输出词概率分布
5. 推理与训练的区别
- 概念说明:推理是逐步生成,训练是并行计算。
- 关键细节:
- 推理:基于前一步的预测结果作为下一步输入
- 训练:基于正确标签并行计算所有位置的 Loss
🧮 核心公式与推导
无(本节为流程描述)
🏗️ 模型架构与数据流向
💻 代码实战
import torch
import torch.nn as nn
import torch.nn.functional as F
class TransformerInference:
"""Transformer推理过程示例"""
def __init__(self, model, src_vocab, tgt_vocab, max_len=50,
start_token='<S>', end_token='<E>'):
self.model = model
self.max_len = max_len
self.start_token = start_token
self.end_token = end_token
@torch.no_grad()
def translate(self, src, src_mask):
"""
src: [batch, src_len] 源语言token IDs
src_mask: 源语言padding mask
return: [batch, tgt_len] 目标语言token IDs
"""
batch_size = src.size(0)
# 1. 编码器处理
memory = self.model.encode(src, src_mask) # [batch, src_len, d_model]
# 2. 初始化解码器输入: 起始符
ys = torch.ones(batch_size, 1, dtype=torch.long).fill_(self.start_token_id)
for i in range(self.max_len - 1):
# 3. 构建目标mask (因果mask)
tgt_mask = self.make_causal_mask(ys)
# 4. 解码器前向传播
out = self.model.decode(memory, ys, src_mask, tgt_mask)
# out: [batch, tgt_len, d_model]
# 5. 取最后一个位置的输出, 经过线性层+softmax
logits = self.model.generator(out[:, -1, :]) # [batch, vocab_size]
prob = F.softmax(logits, dim=-1)
# 6. 取概率最大的词
next_word = prob.argmax(dim=-1, keepdim=True) # [batch, 1]
# 7. 拼接到解码器输入
ys = torch.cat([ys, next_word], dim=1)
# 8. 遇到结束符则停止
if next_word.item() == self.end_token_id:
break
return ys # [batch, tgt_len]
def make_causal_mask(self, tgt):
"""生成因果掩码 (下三角矩阵)"""
L = tgt.size(1)
mask = torch.tril(torch.ones(L, L)).bool() # 下三角为True
return mask
⚠️ 常见问题与避坑指南
- 推理时每步基于前一步的预测结果,错误会累积(exposure bias)
- 必须设置最大长度
max_len防止无限生成 - 遇到结束符
<E>才停止生成 - 推理是串行的(逐 token 生成),比训练慢
💡 个人总结与延伸
Transformer 的推理过程是典型的自回归生成,解码器基于已生成的 token 逐步预测下一个 token,直到输出结束符。这一机制是所有生成式大模型(GPT 系列)的基础。现代大模型通过 KV Cache、推测解码(Speculative Decoding)等技术加速推理过程。