共计 3640 个字符,预计需要花费 10 分钟才能阅读完成。
传统序列模型的局限性
在 NLP 领域,传统的 RNN 和 LSTM 模型在处理长文本对话时存在几个明显的痛点:

- 并行计算困难:RNN 必须按时间步顺序计算,无法像 Transformer 那样同时处理整个序列
- 长期依赖丢失:即便使用 LSTM,当序列长度超过 100 时,梯度仍然容易消失或爆炸
- 内存占用高:RNN 需要存储所有中间状态,而 Transformer 只需维护注意力权重矩阵
模型架构对比
| 特性 | RNN | LSTM | Transformer |
|---|---|---|---|
| 计算复杂度 | O(n) | O(n) | O(n²) |
| 并行性 | 不支持 | 不支持 | 完全支持 |
| 长期依赖处理 | 差 | 中等 | 优秀 |
| 内存占用 | 中等 | 高 | 可控 |
Transformer 核心实现
1. 基础组件搭建
首先实现最核心的 Multi-Head Attention(多头注意力):
import torch
import torch.nn as nn
import math
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, n_heads=8):
super().__init__()
assert d_model % n_heads == 0
self.d_k = d_model // n_heads
self.n_heads = n_heads
self.q_linear = nn.Linear(d_model, d_model)
self.k_linear = nn.Linear(d_model, d_model)
self.v_linear = nn.Linear(d_model, d_model)
self.out_linear = nn.Linear(d_model, d_model)
def forward(self, q, k, v, mask=None):
"""
输入形状: (batch_size, seq_len, d_model)
输出形状: (batch_size, seq_len, d_model)
"""
batch_size = q.size(0)
# 线性投影 + 分头
q = self.q_linear(q).view(batch_size, -1, self.n_heads, self.d_k)
k = self.k_linear(k).view(batch_size, -1, self.n_heads, self.d_k)
v = self.v_linear(v).view(batch_size, -1, self.n_heads, self.d_k)
# 缩放点积注意力计算
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(self.d_k)
if mask is not None:
scores = scores.masked_fill(mask == 0, -1e9)
attn = torch.softmax(scores, dim=-1)
output = torch.matmul(attn, v)
# 合并多头输出
output = output.transpose(1, 2).contiguous() \
.view(batch_size, -1, self.n_heads * self.d_k)
return self.out_linear(output)
2. 位置编码实现
Transformer 没有时序概念,需要通过 Positional Encoding 注入位置信息:
class PositionalEncoding(nn.Module):
def __init__(self, d_model, max_len=5000):
super().__init__()
pe = torch.zeros(max_len, d_model)
position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)
div_term = torch.exp(torch.arange(0, d_model, 2).float() *
(-math.log(10000.0) / d_model))
pe[:, 0::2] = torch.sin(position * div_term)
pe[:, 1::2] = torch.cos(position * div_term)
self.register_buffer('pe', pe)
def forward(self, x):
return x + self.pe[:x.size(1), :]
对话任务适配
1. 特殊掩码处理
对话生成需要两种 mask:
- Padding Mask:处理变长序列
- Sequence Mask:防止解码器看到未来信息
def create_masks(src, trg):
# src_mask (batch_size, 1, 1, src_len)
src_mask = (src != 0).unsqueeze(1).unsqueeze(2)
# trg_mask (batch_size, 1, trg_len, trg_len)
trg_pad_mask = (trg != 0).unsqueeze(1).unsqueeze(2)
trg_len = trg.shape[1]
trg_sub_mask = torch.tril(torch.ones(trg_len, trg_len)).bool()
trg_mask = trg_pad_mask & trg_sub_mask
return src_mask, trg_mask
2. Beam Search 实现
def beam_search(model, src, beam_size=5, max_len=50):
"""
输入:
src - (1, src_len)
输出:
best_sequence - (1, output_len)
"""
with torch.no_grad():
# 编码器处理
memory = model.encoder(src)
# 初始化 beam
sequences = [[[model.bos_idx], 0.0]]
for _ in range(max_len):
candidates = []
for seq in sequences:
# 解码已生成部分
trg = torch.LongTensor(seq[0]).unsqueeze(0).to(device)
output = model.decoder(trg, memory)
# 取最后一个时间步的 logits
logits = output[:, -1, :]
log_probs = torch.log_softmax(logits, dim=-1)
topk_probs, topk_ids = log_probs.topk(beam_size)
# 扩展候选序列
for i in range(beam_size):
new_seq = seq[0] + [topk_ids[0][i].item()]
new_score = seq[1] + topk_probs[0][i].item()
candidates.append([new_seq, new_score])
# 选择 topk 候选
ordered = sorted(candidates, key=lambda x: x[1], reverse=True)
sequences = ordered[:beam_size]
# 检查是否全部生成 EOS
if all([seq[0][-1] == model.eos_idx for seq in sequences]):
break
return sequences[0][0]
生产环境优化
1. 混合精度训练
from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
for batch in dataloader:
optimizer.zero_grad()
with autocast():
output = model(batch.src, batch.trg)
loss = criterion(output, batch.label)
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
2. 梯度裁剪
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
常见问题排查
- 注意力权重全为 NaN
- 检查是否忘记对注意力分数做缩放(除以 $\sqrt{d_k}$)
-
确认 softmax 维度是否正确(应对最后一个维度做归一化)
-
训练损失不下降
- 检查位置编码是否正常注入
-
验证输入 embedding 是否做了恰当的归一化
-
推理结果重复
- 调整 temperature 参数降低确定性
- 检查 beam search 的宽度是否过小
延伸思考
- 如何利用 KV 缓存 (KV Cache) 优化自回归生成速度?
- 在长对话场景下,如何设计更高效的内存管理机制?
- 对于领域特定的对话任务,预训练 + 微调与传统方法相比有哪些优劣?
通过这个实战项目,我们不仅理解了 Transformer 的核心机制,还掌握了将其应用于对话系统的完整流程。建议读者尝试调整模型深度、注意力头数等超参数,观察对生成质量的影响。
正文完
发表至: 未分类
近一天内
