BERT基础教程:从零开始实现Transformer大模型实战

1次阅读
没有评论

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

image.webp

自然语言处理的序列建模痛点

在自然语言处理(NLP)中,序列建模一直是一个核心挑战。传统方法如 RNN 和 LSTM 虽然能够处理序列数据,但在面对长距离依赖问题时表现不佳。RNN 由于梯度消失问题,难以捕捉相隔较远的词之间的关系;LSTM 虽然通过门控机制缓解了这一问题,但计算效率仍然较低,且难以并行化处理。

BERT 基础教程:从零开始实现 Transformer 大模型实战

Transformer 架构的出现改变了这一局面。它通过自注意力机制(Self-Attention)直接建模序列中任意两个词之间的关系,无论它们相隔多远。这种机制不仅解决了长距离依赖问题,还大幅提升了计算效率,使得模型能够并行处理整个序列。

BERT 核心组件详解

自注意力机制

自注意力机制是 Transformer 的核心,其数学推导如下:

  1. QKV 矩阵:对于输入序列中的每个词,我们生成三个向量——Query(Q)、Key(K)和 Value(V)。这些向量通过线性变换从输入嵌入中获取。

  2. 注意力分数计算:通过计算 Q 和 K 的点积,得到注意力分数,再经过 softmax 归一化,得到权重分布。公式如下:
    $$
    \text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^T}{\sqrt{d_k}}\right)V
    $$
    其中,$d_k$ 是 Key 向量的维度,用于缩放点积,防止梯度消失。

  3. 多头注意力:为了捕捉不同子空间的信息,BERT 使用多头注意力机制,即将 Q、K、V 分成多组,分别计算注意力后拼接起来。

位置编码

由于 Transformer 没有递归结构,需要显式地注入位置信息。BERT 使用三角函数生成位置编码:

$$
PE_{(pos, 2i)} = \sin\left(\frac{pos}{10000^{2i/d_{\text{model}}}}\right)
$$
$$
PE_{(pos, 2i+1)} = \cos\left(\frac{pos}{10000^{2i/d_{\text{model}}}}\right)
$$

其中,$pos$ 是位置,$i$ 是维度索引,$d_{\text{model}}$ 是模型维度。这种编码方式能够捕捉相对位置信息,并且对长序列具有良好的泛化能力。

层归一化与残差连接

层归一化(Layer Normalization)和残差连接(Residual Connection)是 BERT 训练稳定的关键:

  • 层归一化:对每个子层的输出进行归一化,缓解梯度消失问题。
  • 残差连接:将输入直接加到子层输出上,确保梯度能够直接回传,加速训练。

PyTorch 实现

以下是一个模块化的 BERT 实现,包含 Embedding、TransformerLayer 和 BERT 类:

import torch
import torch.nn as nn
import math

class Embedding(nn.Module):
    def __init__(self, vocab_size, d_model, max_len):
        super().__init__()
        self.token_embed = nn.Embedding(vocab_size, d_model)
        self.pos_embed = nn.Embedding(max_len, d_model)
        self.norm = nn.LayerNorm(d_model)

    def forward(self, x):
        pos = torch.arange(x.size(1), device=x.device).unsqueeze(0)
        x = self.token_embed(x) + self.pos_embed(pos)
        return self.norm(x)

class TransformerLayer(nn.Module):
    def __init__(self, d_model, n_head):
        super().__init__()
        self.attention = nn.MultiheadAttention(d_model, n_head)
        self.ffn = nn.Sequential(nn.Linear(d_model, 4 * d_model),
            nn.GELU(),
            nn.Linear(4 * d_model, d_model)
        )
        self.norm1 = nn.LayerNorm(d_model)
        self.norm2 = nn.LayerNorm(d_model)

    def forward(self, x, mask=None):
        # 自注意力
        attn_out, _ = self.attention(x, x, x, key_padding_mask=mask)
        x = x + attn_out
        x = self.norm1(x)
        # 前馈网络
        ffn_out = self.ffn(x)
        x = x + ffn_out
        x = self.norm2(x)
        return x

class BERT(nn.Module):
    def __init__(self, vocab_size, d_model, n_head, n_layers, max_len):
        super().__init__()
        self.embed = Embedding(vocab_size, d_model, max_len)
        self.layers = nn.ModuleList([TransformerLayer(d_model, n_head) for _ in range(n_layers)])
        self.classifier = nn.Linear(d_model, 2)  # 假设是二分类任务

    def forward(self, x, mask=None):
        x = self.embed(x)
        for layer in self.layers:
            x = layer(x, mask)
        # 使用 [CLS] 向量进行分类
        cls_out = x[:, 0, :]
        return self.classifier(cls_out)

生产环境注意事项

长序列处理

当输入序列超过 512token 时,可以采用以下方案:

  1. 截断:直接截断超过 512 的部分。
  2. 滑动窗口:将长序列分成多个 512token 的片段,分别处理后合并结果。
  3. 稀疏注意力:使用稀疏注意力机制(如 Longformer)减少计算量。

混合精度训练

使用混合精度训练(FP16)可以大幅减少内存占用:

from torch.cuda.amp import autocast, GradScaler

scaler = GradScaler()

with autocast():
    outputs = model(inputs)
    loss = criterion(outputs, labels)

scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()

分布式训练

使用 HuggingFace Accelerate 可以轻松实现分布式训练:

from accelerate import Accelerator

accelerator = Accelerator()
model, optimizer, train_loader = accelerator.prepare(model, optimizer, train_loader)

for batch in train_loader:
    optimizer.zero_grad()
    outputs = model(batch['input_ids'])
    loss = criterion(outputs, batch['labels'])
    accelerator.backward(loss)
    optimizer.step()

思考题

  1. 如何修改架构实现中文 BERT?
  2. 中文 BERT 通常使用字或词作为输入单元,需要调整分词器和词汇表。

  3. 对比 BERT 与 GPT 的位置编码差异

  4. BERT 使用绝对位置编码,而 GPT 使用相对位置编码。

  5. 解释 [CLS] 向量为何能用于分类任务

  6. [CLS]向量在预训练时被设计为捕捉整个序列的全局信息,因此适合用于分类任务。

总结

本文从零开始实现了 BERT 模型的核心组件,并提供了完整的 PyTorch 代码。通过自注意力机制、位置编码和层归一化等技术,BERT 能够高效处理序列数据。在生产环境中,我们还探讨了长序列处理、混合精度训练和分布式训练的优化技巧。希望这篇教程能帮助你更好地理解和使用 BERT 模型。

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