从零开始构建Transformer模型:原理详解与PyTorch实战指南

1次阅读
没有评论

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

image.webp

Transformer 的革命性意义

Transformer 彻底改变了 NLP 领域处理序列数据的方式,完全摒弃了传统的循环神经网络结构。它通过自注意力机制实现了对长距离依赖的高效建模,使得并行计算成为可能。最重要的是,这种架构为后来的 BERT、GPT 等预训练模型奠定了基础,成为现代 NLP 的核心组件。

从零开始构建 Transformer 模型:原理详解与 PyTorch 实战指南

模型构建详解

1. 词嵌入与位置编码

Transformer 首先需要将单词转换为固定维度的向量表示。这里我们使用 PyTorch 的 nn.Embedding 来实现词嵌入层:

import torch.nn as nn
vocab_size = 10000  # 词汇表大小
d_model = 512       # 嵌入维度
embedding = nn.Embedding(vocab_size, d_model)

由于 Transformer 没有循环结构,我们需要显式地加入位置信息。下面是正弦位置编码的实现:

import torch
import math

class PositionalEncoding(nn.Module):
    def __init__(self, d_model, max_len=5000):
        super().__init__()
        position = torch.arange(max_len).unsqueeze(1)
        div_term = torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model))
        pe = torch.zeros(max_len, 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)]

2. 自注意力机制实现

自注意力的核心是 QKV(Query-Key-Value)三元组计算。我们先定义线性变换层:

self.query = nn.Linear(d_model, d_model)
self.key = nn.Linear(d_model, d_model)
self.value = nn.Linear(d_model, d_model)

然后实现缩放点积注意力:

def scaled_dot_product_attention(q, k, v, mask=None):
    matmul_qk = torch.matmul(q, k.transpose(-2, -1))
    d_k = q.size(-1)
    scaled_attention_logits = matmul_qk / math.sqrt(d_k)

    if mask is not None:
        scaled_attention_logits += (mask * -1e9)

    attention_weights = torch.softmax(scaled_attention_logits, dim=-1)
    output = torch.matmul(attention_weights, v)
    return output

3. 多头注意力并行化

将注意力头拆分成多个子空间可以捕捉不同的语义信息。关键实现技巧是使用 viewtranspose操作:

class MultiHeadAttention(nn.Module):
    def __init__(self, d_model, num_heads):
        super().__init__()
        self.num_heads = num_heads
        self.d_model = d_model
        assert d_model % num_heads == 0
        self.depth = d_model // num_heads

        self.wq = nn.Linear(d_model, d_model)
        self.wk = nn.Linear(d_model, d_model)
        self.wv = nn.Linear(d_model, d_model)
        self.dense = nn.Linear(d_model, d_model)

    def split_heads(self, x, batch_size):
        x = x.view(batch_size, -1, self.num_heads, self.depth)
        return x.transpose(1, 2)

    def forward(self, q, k, v, mask=None):
        batch_size = q.size(0)
        q = self.split_heads(self.wq(q), batch_size)
        k = self.split_heads(self.wk(k), batch_size)
        v = self.split_heads(self.wv(v), batch_size)

        scaled_attention = scaled_dot_product_attention(q, k, v, mask)
        concat_attention = scaled_attention.transpose(1, 2).contiguous()
        concat_attention = concat_attention.view(batch_size, -1, self.d_model)
        output = self.dense(concat_attention)
        return output

完整文本分类实现

数据准备

我们使用 HuggingFace 的 IMDB 数据集:

from datasets import load_dataset
imdb = load_dataset('imdb')
tokenizer = ... # 实现分词器

模型训练技巧

为了防止训练初期的不稳定,我们实现学习率预热和梯度裁剪:

optimizer = torch.optim.Adam(model.parameters(), lr=0, betas=(0.9, 0.98), eps=1e-9)

# 学习率预热
lr_scheduler = torch.optim.lr_scheduler.LambdaLR(
    optimizer,
    lambda step: min((step+1)**-0.5, (step+1)*warmup_steps**-1.5)
)

# 梯度裁剪
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

避坑指南

数值稳定性

计算注意力分数时,较大的输入会导致 softmax 进入梯度饱和区。解决方案:

scaled_attention_logits = matmul_qk / math.sqrt(d_k)

显存管理

Transformer 的显存占用主要来自注意力矩阵,其空间复杂度为 O(n²)。估算公式:

显存(B) ≈ batch_size × seq_len² × num_heads × 4 (float32)

对于长文本,可以尝试:
– 减小 batch_size
– 使用梯度累积
– 采用线性注意力变体

进阶思考

  1. 如何修改架构实现字符级 Transformer?
  2. 将词嵌入层替换为字符卷积
  3. 调整位置编码的最大长度

  4. 对比 CNN 在长文本处理的优劣:

  5. CNN 的局部感受野限制长距离建模
  6. Transformer 的自注意力具有全局视野
  7. CNN 在短文本上可能有计算效率优势

  8. 解释 LayerNorm 对梯度传播的影响:

  9. 保持各层输入的分布稳定
  10. 缓解梯度消失 / 爆炸问题
  11. 与 BatchNorm 相比更适合变长输入

结语

通过本文的实践,你应该已经掌握了 Transformer 的核心实现要点。建议在实际项目中先从小规模数据开始实验,逐步调整模型规模和训练策略。Transformer 虽然强大,但也需要根据具体任务进行适当的调整和优化。

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