ChatGPT中的Transformer架构实战:从零搭建你的第一个对话模型

1次阅读
没有评论

共计 3640 个字符,预计需要花费 10 分钟才能阅读完成。

image.webp

传统序列模型的局限性

在 NLP 领域,传统的 RNN 和 LSTM 模型在处理长文本对话时存在几个明显的痛点:

ChatGPT 中的 Transformer 架构实战:从零搭建你的第一个对话模型

  • 并行计算困难: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)

常见问题排查

  1. 注意力权重全为 NaN
  2. 检查是否忘记对注意力分数做缩放(除以 $\sqrt{d_k}$)
  3. 确认 softmax 维度是否正确(应对最后一个维度做归一化)

  4. 训练损失不下降

  5. 检查位置编码是否正常注入
  6. 验证输入 embedding 是否做了恰当的归一化

  7. 推理结果重复

  8. 调整 temperature 参数降低确定性
  9. 检查 beam search 的宽度是否过小

延伸思考

  • 如何利用 KV 缓存 (KV Cache) 优化自回归生成速度?
  • 在长对话场景下,如何设计更高效的内存管理机制?
  • 对于领域特定的对话任务,预训练 + 微调与传统方法相比有哪些优劣?

通过这个实战项目,我们不仅理解了 Transformer 的核心机制,还掌握了将其应用于对话系统的完整流程。建议读者尝试调整模型深度、注意力头数等超参数,观察对生成质量的影响。

正文完
 0
评论(没有评论)