共计 3146 个字符,预计需要花费 8 分钟才能阅读完成。
背景与核心概念
BERT(Bidirectional Encoder Representations from Transformers)是自然语言处理(NLP)领域的重要模型,其预训练过程依赖于有效的输入表示。标记嵌入(Token Embedding)和位置编码(Positional Encoding)是 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}})$$
这种编码方式具有以下优势:
- 可以表示任意长度的序列
- 能捕捉相对位置关系
- 允许模型轻松学习关注相对位置
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)
性能考量与优化
- 矩阵运算优化:
- 使用 PyTorch 的
torch.bmm进行批量矩阵乘法 -
启用 CUDA 加速和
torch.backends.cudnn.benchmark = True -
缓存机制:
- 预计算位置编码并缓存
-
对静态词表使用 EmbeddingBag
-
混合精度训练:
- 使用
torch.cuda.amp自动混合精度 -
减少显存占用同时保持精度
-
批处理策略:
- 动态 padding 和 masking
- 相似长度样本批处理
避坑指南
- 词表不匹配:
- 问题:训练和推理使用的词表不一致
-
解决:确保使用相同的 tokenizer 和词表文件
-
位置编码溢出:
- 问题:序列长度超过预定义的 max_len
-
解决:动态扩展位置编码或截断长序列
-
梯度爆炸:
- 问题:嵌入层梯度异常增大
-
解决:添加梯度裁剪(gradient clipping)
-
显存不足:
- 问题:长序列导致 OOM
-
解决:减小 batch size 或使用梯度累积
-
数值不稳定:
- 问题:位置编码值过大或过小
- 解决:对嵌入进行 LayerNorm
总结与思考题
本文详细介绍了 BERT 预训练中标记嵌入和位置编码的实现方法。通过 PyTorch 代码示例,展示了如何构建完整的输入表示层。关键要点包括:
- WordPiece 标记化的优势
- 正弦位置编码的数学原理
- 实际实现中的性能优化
思考:对于中文文本,位置编码策略可以做哪些调整?
- 考虑中文分词特性是否影响位置编码效果
- 中文长文本是否需要调整位置编码的频率参数
- 如何处理中文中常见的标点符号密集情况
欢迎在评论区分享你的想法和实践经验!
