BERT模型预训练实战:从标记嵌入层到位置编码的完整实现

1次阅读
没有评论

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

image.webp

背景与核心概念

BERT(Bidirectional Encoder Representations from Transformers)是自然语言处理(NLP)领域的重要模型,其预训练过程依赖于有效的输入表示。标记嵌入(Token Embedding)和位置编码(Positional Encoding)是 BERT 输入表示的两大核心组件。

BERT 模型预训练实战:从标记嵌入层到位置编码的完整实现

标记嵌入负责将离散的文本符号(如单词或子词)映射到连续的向量空间,而位置编码则注入序列的顺序信息。这两种嵌入的结合使模型能够同时理解语义和序列关系。

标记嵌入层实现

标记化与嵌入方式对比

BERT 采用 WordPiece 分词算法,与传统的 BPE(Byte Pair Encoding)相比有以下特点:

  • WordPiece 基于概率合并子词,而 BPE 基于频率
  • WordPiece 能更好处理罕见词,适合多语言场景
  • BPE 实现更简单,但 WordPiece 与 BERT 配合更优

PyTorch 实现

import torch
import torch.nn as nn

class TokenEmbedding(nn.Module):
    """
    WordPiece 标记嵌入层
    参数:
        vocab_size: 词表大小
        hidden_size: 嵌入维度
    """
    def __init__(self, vocab_size: int, hidden_size: int):
        super().__init__()
        self.embedding = nn.Embedding(vocab_size, hidden_size)
        self.hidden_size = hidden_size

    def forward(self, input_ids: torch.Tensor) -> torch.Tensor:
        """
        输入:
            input_ids: [batch_size, seq_len]
        输出:
            embeddings: [batch_size, seq_len, hidden_size]
        """
        return self.embedding(input_ids) * (self.hidden_size ** 0.5)

位置编码原理与实现

正弦函数编码原理

位置编码使用不同频率的正弦和余弦函数组合:

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

这种编码方式具有以下优势:

  1. 可以表示任意长度的序列
  2. 能捕捉相对位置关系
  3. 允许模型轻松学习关注相对位置

PyTorch 实现

class PositionalEncoding(nn.Module):
    """
    正弦位置编码层
    参数:
        max_len: 最大序列长度
        hidden_size: 嵌入维度
        dropout: Dropout 概率
    """
    def __init__(self, max_len: int, hidden_size: int, dropout: float = 0.1):
        super().__init__()
        self.dropout = nn.Dropout(p=dropout)

        position = torch.arange(max_len).unsqueeze(1)
        div_term = torch.exp(torch.arange(0, hidden_size, 2) * (-math.log(10000.0) / hidden_size))

        pe = torch.zeros(max_len, hidden_size)
        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: torch.Tensor) -> torch.Tensor:
        """
        输入:
            x: [batch_size, seq_len, hidden_size]
        输出:
            x + positional_encoding
        """
        x = x + self.pe[:x.size(1)]
        return self.dropout(x)

完整代码示例

将标记嵌入和位置编码组合使用:

import math

class BERTEmbedding(nn.Module):
    """
    BERT 输入表示层
    组合标记嵌入、位置编码和段嵌入
    """
    def __init__(self, vocab_size: int, hidden_size: int, max_len: int, dropout: float = 0.1):
        super().__init__()
        self.token_embedding = TokenEmbedding(vocab_size, hidden_size)
        self.position_encoding = PositionalEncoding(max_len, hidden_size, dropout)

    def forward(self, input_ids: torch.Tensor) -> torch.Tensor:
        """
        输入:
            input_ids: [batch_size, seq_len]
        输出:
            embeddings: [batch_size, seq_len, hidden_size]
        """
        token_embeddings = self.token_embedding(input_ids)
        embeddings = self.position_encoding(token_embeddings)
        return embeddings

处理示例句子:

# 假设词表已包含这些词
sentences = ["paris is a beautiful city", "i love paris"]

# 创建模型实例
embedder = BERTEmbedding(
    vocab_size=30000,
    hidden_size=768,
    max_len=512
)

# 将句子转换为 token ids (简化示例)
input_ids = torch.tensor([[12, 8, 3, 45, 23],  # "paris is a beautiful city"
    [7, 90, 12]          # "i love paris"
])

# 获取嵌入表示
embeddings = embedder(input_ids)

性能考量与优化

  1. 矩阵运算优化
  2. 使用 PyTorch 的 torch.bmm 进行批量矩阵乘法
  3. 启用 CUDA 加速和torch.backends.cudnn.benchmark = True

  4. 缓存机制

  5. 预计算位置编码并缓存
  6. 对静态词表使用 EmbeddingBag

  7. 混合精度训练

  8. 使用 torch.cuda.amp 自动混合精度
  9. 减少显存占用同时保持精度

  10. 批处理策略

  11. 动态 padding 和 masking
  12. 相似长度样本批处理

避坑指南

  1. 词表不匹配
  2. 问题:训练和推理使用的词表不一致
  3. 解决:确保使用相同的 tokenizer 和词表文件

  4. 位置编码溢出

  5. 问题:序列长度超过预定义的 max_len
  6. 解决:动态扩展位置编码或截断长序列

  7. 梯度爆炸

  8. 问题:嵌入层梯度异常增大
  9. 解决:添加梯度裁剪(gradient clipping)

  10. 显存不足

  11. 问题:长序列导致 OOM
  12. 解决:减小 batch size 或使用梯度累积

  13. 数值不稳定

  14. 问题:位置编码值过大或过小
  15. 解决:对嵌入进行 LayerNorm

总结与思考题

本文详细介绍了 BERT 预训练中标记嵌入和位置编码的实现方法。通过 PyTorch 代码示例,展示了如何构建完整的输入表示层。关键要点包括:

  • WordPiece 标记化的优势
  • 正弦位置编码的数学原理
  • 实际实现中的性能优化

思考:对于中文文本,位置编码策略可以做哪些调整?

  1. 考虑中文分词特性是否影响位置编码效果
  2. 中文长文本是否需要调整位置编码的频率参数
  3. 如何处理中文中常见的标点符号密集情况

欢迎在评论区分享你的想法和实践经验!

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