🎯 课程主题
Transformer 模型的推理与应用——加载训练好的权重进行翻译。
📝 核心知识点
1. 推理流程
- 概念说明:自回归生成翻译结果。
- 关键细节:
- 加载训练好的权重
- 输入源语言句子(如英文)
- 编码器处理源语言 → 编码器信息
- 解码器从起始符
<S>开始生成 - 逐词生成,直到结束符
</s> - 输出目标语言句子(如中文)
2. 自回归生成
- 概念说明:每步生成一个词,将生成的词加入输入。
- 关键细节:
- 初始输入:
<S>(起始符) - 第 1 步:输入
<S>→ 生成第 1 个词 - 第 2 步:输入
<S> 词1→ 生成第 2 个词 - 第 N 步:输入
<S> 词1...词N-1→ 生成第 N 个词 - 直到生成
</s>(结束符)或达到最大长度
- 初始输入:
3. 权重配置
- 概念说明:推理时需要配置权重路径。
- 关键细节:
- 在配置文件中设置权重路径
- 可使用预训练权重或自己训练的权重
- 推理时只需前向传播,无需反向传播
4. 翻译效果
- 概念说明:模型翻译结果展示。
- 关键细节:
- "i love you" → "我爱你"
- "i am a dog" → "我是个狗"
- "The government has implemented a series of policies to improve the living standards of citizens" → "政府实施了一系列政策改善国民生活水平"
- BLEU 26.3 的翻译质量已能正确表达意思
5. 性能提升方向
- 概念说明:如何提升模型性能。
- 关键细节:
- 增加训练数据
- 增大模型(更多层、更大维度)
- 调整超参数
- 使用更先进的架构(如大模型技术)
🧮 核心公式与推导
自回归生成:
$$y_1 = \text{Decode}(\text{Enc}(src), \langle S \rangle)$$
$$y_2 = \text{Decode}(\text{Enc}(src), \langle S \rangle, y_1)$$
$$y_t = \text{Decode}(\text{Enc}(src), \langle S \rangle, y_1, ..., y_{t-1})$$
直到 $y_t = \langle /s \rangle$ 或 $t = T_{max}$
每步预测:
$$y_t = \arg\max_{i} P(y_{t,i} | src, y_{<t})$$
贪心解码:每步取概率最大的词
🏗️ 模型架构与数据流向
💻 代码实战
import torch
from config import Config
from model import make_model
from tokenizer import Tokenizer
class Translator:
"""翻译器: 加载模型进行推理"""
def __init__(self, config):
self.config = config
self.device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
# 1. 加载分词器
self.src_tokenizer = Tokenizer(config.src_tokenizer_path)
self.tgt_tokenizer = Tokenizer(config.tgt_tokenizer_path)
# 2. 构建模型
self.model = make_model(
src_vocab=config.src_vocab_size,
tgt_vocab=config.tgt_vocab_size,
N=config.n_layers,
d_model=config.d_model,
d_ff=config.d_ff,
h=config.n_heads,
dropout=0 # 推理时关闭dropout
).to(self.device)
# 3. 加载权重
checkpoint = torch.load(config.model_path, map_location=self.device)
self.model.load_state_dict(checkpoint)
self.model.eval() # 推理模式
@torch.no_grad()
def translate(self, text, max_len=50):
"""
翻译函数 (贪心解码)
text: 源语言文本
max_len: 最大生成长度
return: 目标语言文本
"""
# 1. 编码源语言
src_ids = self.src_tokenizer.encode(text)
src_tensor = torch.LongTensor([src_ids]).to(self.device)
src_mask = self.make_src_mask(src_tensor)
# 2. 编码器前向传播 (只计算一次)
memory = self.model.encode(src_tensor, src_mask)
# 3. 自回归生成
# 初始输入: <S> (起始符)
tgt_ids = [self.config.bos_idx]
for _ in range(max_len):
tgt_tensor = torch.LongTensor([tgt_ids]).to(self.device)
tgt_mask = self.make_tgt_mask(tgt_tensor)
# 解码器前向传播
output = self.model.decode(
memory, src_mask, tgt_tensor, tgt_mask
)
# 生成器: 取最后一个位置的预测
logits = self.model.generator(output[:, -1, :])
next_token = logits.argmax(dim=-1).item()
# 添加到序列
tgt_ids.append(next_token)
# 遇到结束符停止
if next_token == self.config.eos_idx:
break
# 4. 解码为目标语言文本
# 去掉起始符和结束符
result_ids = tgt_ids[1:-1] if tgt_ids[-1] == self.config.eos_idx else tgt_ids[1:]
result = self.tgt_tokenizer.decode(result_ids)
return result
def make_src_mask(self, src):
"""源语言填充掩码"""
return (src != self.config.pad_idx).unsqueeze(-2)
def make_tgt_mask(self, tgt):
"""目标语言掩码: 填充掩码 + 因果掩码"""
tgt_pad_mask = (tgt != self.config.pad_idx).unsqueeze(-2)
tgt_len = tgt.size(1)
# 因果掩码 (下三角)
tgt_causal_mask = torch.tril(
torch.ones(tgt_len, tgt_len, device=self.device)
).bool()
# 合并两个掩码
return tgt_pad_mask & tgt_causal_mask
# 使用示例
if __name__ == "__main__":
config = Config()
config.model_path = 'runs/best.pt' # 权重路径
translator = Translator(config)
# 测试翻译
test_sentences = [
"i love you",
"i am a dog",
"The government has implemented a series of policies to improve the living standards of citizens"
]
for src in test_sentences:
result = translator.translate(src)
print(f"英文: {src}")
print(f"中文: {result}")
print("-" * 50)
⚠️ 常见问题与避坑指南
- 推理时
model.eval()关闭 Dropout - 使用
@torch.no_grad()不计算梯度,节省内存 - 编码器只需计算一次,解码器逐步生成
- 设置
max_len防止无限生成 - 推理时需要同时使用填充掩码和因果掩码
- 贪心解码可能不是最优,可用 Beam Search 提升质量
💡 个人总结与延伸
模型推理是 Transformer 应用的最终环节。自回归生成是 Decoder 的核心特性,通过逐步生成实现序列建模。现代大模型推理优化包括 KV Cache(缓存键值对避免重复计算)、Beam Search(提升生成质量)、量化(减少内存占用)等技术。本课程通过完整的训练-推理流程,展示了 Transformer 在机器翻译任务中的实际应用,为学习更复杂的大模型打下坚实基础。
课程总结: 本课程从 Transformer 的原理到实战,完整覆盖了:
- 理论基础:注意力机制、位置编码、多头注意力、层归一化、FFN 等
- 架构理解:编码器-解码器结构、三种注意力机制、残差连接
- 代码实现:从基础模块到完整模型,PyTorch 实现
- 训练与推理:数据处理、词表构建、模型训练、翻译应用
通过本课程的学习,可以深入理解 Transformer 的原理和实现,为学习 BERT、GPT 等现代大模型奠定基础。