共计 3307 个字符,预计需要花费 9 分钟才能阅读完成。
背景介绍
BERT(Bidirectional Encoder Representations from Transformers)是自然语言处理(NLP)领域的一项里程碑式技术。其核心创新之一就是掩码语言模型(Masked Language Model, MLM)。MLM 通过随机遮盖输入文本中的部分单词,让模型预测这些被遮盖的单词,从而学习到上下文相关的词表示。这种方法使 BERT 能够捕获双向上下文信息,大大提升了在各种 NLP 任务上的表现。

核心公式解析
1. Softmax 交叉熵损失函数
MLM 的核心是一个分类任务,即预测被遮盖单词的正确词汇。这通常通过 softmax 交叉熵损失函数来实现。具体公式如下:
$$
\mathcal{L}{MLM} = -\sum y_i \log(p_i)
$$}^{|V|
其中,$|V|$ 是词汇表大小,$y_i$ 是真实标签的 one-hot 编码,$p_i$ 是模型预测的概率分布:
$$
p_i = \frac{\exp(z_i)}{\sum_{j=1}^{|V|} \exp(z_j)}
$$
这里,$z_i$ 是模型输出的 logits。直观上,这个损失函数衡量的是模型预测分布与真实分布之间的差异。
2. 注意力机制的参与
BERT 使用 Transformer 的多头自注意力机制来捕获上下文信息。对于每个被遮盖的 token,模型通过自注意力机制聚合其周围 token 的信息。具体来说,对于第 $l$ 层的第 $h$ 个头,注意力权重计算如下:
$$
\text{Attention}(Q, K, V) = \text{softmax}(\frac{QK^T}{\sqrt{d_k}})V
$$
其中,$Q$、$K$、$V$ 分别是通过线性变换得到的查询、键和值矩阵,$d_k$ 是键向量的维度。通过多层多头注意力,BERT 能够从不同角度和层次捕获上下文信息。
PyTorch 实现示例
以下是一个简化的 MLM 实现示例,展示了关键部分的代码:
import torch
import torch.nn as nn
import torch.nn.functional as F
class BERTMLM(nn.Module):
def __init__(self, vocab_size, hidden_size):
super().__init__()
self.embedding = nn.Embedding(vocab_size, hidden_size)
self.transformer = nn.TransformerEncoder(nn.TransformerEncoderLayer(hidden_size, nhead=8),
num_layers=6
)
self.fc = nn.Linear(hidden_size, vocab_size)
def forward(self, input_ids, attention_mask=None):
# input_ids: [batch_size, seq_len]
embeds = self.embedding(input_ids) # [batch_size, seq_len, hidden_size]
# Transformer 编码
outputs = self.transformer(embeds.transpose(0, 1), # Transformer 需要 seq_len 在前
src_key_padding_mask=attention_mask
).transpose(0, 1)
# 预测被 mask 的 token
logits = self.fc(outputs) # [batch_size, seq_len, vocab_size]
return logits
# 使用示例
model = BERTMLM(vocab_size=30000, hidden_size=768)
criterion = nn.CrossEntropyLoss(ignore_index=-100) # 忽略非 mask 位置的损失
# 假设输入数据
input_ids = torch.randint(0, 30000, (32, 128)) # 模拟 batch_size=32, seq_len=128
masked_positions = torch.randint(0, 128, (32, 20)) # 每个样本 mask 20 个位置
labels = torch.randint(0, 30000, (32, 128)) # 真实标签
labels[~masked_positions] = -100 # 只计算 mask 位置的损失
logits = model(input_ids)
loss = criterion(logits.view(-1, 30000), labels.view(-1))
关键优化点:
- 使用
ignore_index避免计算非 mask 位置的损失,提高训练效率 - 合理设置 batch_size 和序列长度以充分利用 GPU 显存
- 在 Transformer 层使用
src_key_padding_mask处理变长序列
工程实践
1. 大规模训练显存优化
当处理大规模语料时,显存成为主要瓶颈。以下是一些优化技巧:
- 使用梯度累积(gradient accumulation):通过多次前向传播累积梯度,再一次性更新参数,可以模拟更大的 batch size 而不会增加显存占用
- 激活值检查点(activation checkpointing):只保存部分层的激活值,需要时重新计算,以空间换时间
- 序列截断和动态填充:根据实际长度动态 padding,避免使用固定最大长度
2. 混合精度训练
混合精度训练可以显著减少显存占用并加速训练:
from torch.cuda.amp import autocast, GradScaler
scaler = GradScaler()
with autocast():
logits = model(input_ids)
loss = criterion(logits.view(-1, 30000), labels.view(-1))
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
注意事项:
- 适当调整
scaler的初始值(通常为 65536.0) - 监控梯度缩放因子,避免下溢或上溢
- 某些操作(如 softmax)可能需要保持 fp32 精度
3. 分布式训练
对于超大规模训练,分布式是必选项:
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
# 初始化分布式环境
dist.init_process_group(backend='nccl')
# 包装模型
model = DDP(model.to(device), device_ids=[local_rank])
# 数据并行采样器
train_sampler = DistributedSampler(dataset)
关键点:
- 确保每进程看到不同的数据分片
- 注意梯度同步的开销
- 合理设置通信后端(通常 NCCL 对 GPU 最有效)
避坑指南
常见实现错误
- mask 比例不当:通常 15% 的 mask 比例效果最佳,过高会导致模型难以学习,过低则训练不充分
- 静态 mask:应在每个 epoch 重新随机 mask,避免模型记忆特定模式
- 忽略特殊 token:不应 mask [CLS]、[SEP]等特殊 token
超参数调优建议
- 学习率:通常 3e- 5 到 5e- 5 之间,使用线性 warmup
- batch size:尽可能大(受显存限制),但要注意梯度噪声的影响
- 层数和头数:base 模型常用 12 层 12 头,large 模型 24 层 16 头
性能考量
硬件对比
| 硬件 | 吞吐量(tokens/sec) | 显存占用 |
|---|---|---|
| V100 32GB | 12k | 28GB |
| A100 40GB | 25k | 36GB |
| T4 16GB | 5k | 14GB |
批大小影响
- 过小:训练不稳定,梯度噪声大
- 过大:可能收敛到 sharp minima,泛化性差
- 建议:在显存允许下尽可能大,通常 256-1024 范围
开放性问题
- 如何设计更有效的 mask 策略(如 span masking)来提升模型性能?
- 在大规模预训练中,如何平衡计算效率和模型表现?
- MLM 损失与其他辅助损失(如 NSP)如何协同优化?
