共计 2720 个字符,预计需要花费 7 分钟才能阅读完成。
背景介绍
BERT(Bidirectional Encoder Representations from Transformers)是自然语言处理(NLP)领域的重要里程碑。其中,掩码语言模型(Masked Language Model, MLM)是 BERT 预训练阶段的核心任务之一。MLM 通过随机掩盖输入文本中的部分词汇,要求模型预测被掩盖的词,从而使模型能够学习到丰富的上下文信息。这种方法突破了传统语言模型只能单向预测的局限,实现了真正的双向上下文理解。

MLM 的重要性在于:
- 它迫使模型理解每个词与其上下文的关系,而不仅仅是记住词语的共现模式
- 通过预测被掩盖的词,模型学习到了丰富的语义和语法知识
- 这种预训练方式可以迁移到各种下游任务,显著提升模型性能
核心公式解析
softmax 交叉熵损失函数推导
MLM 的核心是 softmax 交叉熵损失函数。对于给定的被掩盖位置,模型需要预测正确的词。具体公式如下:
- 首先计算每个词的概率分布:
$$P(w_i|C) = \frac{\exp(h^T \cdot E_{w_i})}{\sum_{j=1}^V \exp(h^T \cdot E_{w_j})}$$
其中,$h$ 是被掩盖位置的隐藏状态,$E_{w_i}$ 是词 $w_i$ 的嵌入向量,$V$ 是词汇表大小。
- 然后计算交叉熵损失:
$$L = -\sum_{i=1}^V y_i \log(P(w_i|C))$$
其中,$y_i$ 是 one-hot 编码的真实标签。
掩码机制对梯度传播的影响
掩码机制不仅影响输入,还会影响梯度传播:
- 只有被掩盖的位置会产生梯度
- 未被掩盖的位置梯度为零
- 这种设计使得模型专注于学习预测任务,而不是简单记忆输入
可视化词向量空间
训练过程中,词向量空间会发生有趣的变化:
- 语义相似的词会在向量空间中聚集
- 语法功能相似的词会形成特定模式
- 多义词的不同含义会分布在不同的子空间中
PyTorch 代码实现
下面是一个完整的 MLM 任务实现示例:
import torch
import torch.nn as nn
import torch.nn.functional as F
class MLMHead(nn.Module):
def __init__(self, hidden_size, vocab_size):
super().__init__()
self.dense = nn.Linear(hidden_size, hidden_size)
self.layer_norm = nn.LayerNorm(hidden_size)
self.decoder = nn.Linear(hidden_size, vocab_size)
def forward(self, hidden_states):
# hidden_states: [batch_size, seq_len, hidden_size]
x = self.dense(hidden_states)
x = F.gelu(x)
x = self.layer_norm(x)
logits = self.decoder(x) # [batch_size, seq_len, vocab_size]
return logits
# 示例使用
batch_size = 32
seq_len = 128
hidden_size = 768
vocab_size = 30522 # BERT 的词汇表大小
# 模拟输入
hidden_states = torch.randn(batch_size, seq_len, hidden_size)
mlm_head = MLMHead(hidden_size, vocab_size)
# 计算 logits
logits = mlm_head(hidden_states)
print(f"Logits shape: {logits.shape}")
# 计算损失
masked_positions = torch.randint(0, seq_len, (batch_size,)) # 随机掩盖位置
true_labels = torch.randint(0, vocab_size, (batch_size,)) # 真实标签
# 只计算被掩盖位置的损失
loss = F.cross_entropy(logits[torch.arange(batch_size), masked_positions],
true_labels
)
print(f"MLM loss: {loss.item():.4f}")
关键点说明:
MLMHead结构包含一个全连接层、GELU 激活函数、LayerNorm 和最后的输出层- 损失计算时只考虑被掩盖的位置,忽略其他位置
- 使用交叉熵损失函数衡量预测与真实标签的差异
实践建议
学习率与 batch size 设置
- 初始学习率建议在 1e- 5 到 5e- 5 之间
- batch size 根据 GPU 显存选择,通常 256-1024 效果较好
- 使用线性 warmup 策略,前 10% 的训练步数逐渐增加学习率
长文本处理技巧
- 对于超过 512token 的文本,可以采用以下策略:
- 滑动窗口分割
- 随机截取片段
- 关键句子选择
- 注意保持句子的完整性,避免在单词中间截断
常见训练陷阱及解决方案
- 损失不下降:
- 检查数据预处理是否正确
- 验证梯度是否正常传播
-
尝试降低学习率
-
过拟合:
- 增加 dropout 比率
- 使用更小的模型
-
添加更多的训练数据
-
训练不稳定:
- 使用梯度裁剪
- 调整 warmup 步数
- 尝试不同的优化器
延伸思考
MLM 与其他预训练目标的比较
- MLM vs NSP(下一句预测):
- MLM 学习词汇级知识
- NSP 学习句子间关系
-
现代模型趋向于只使用 MLM
-
MLM vs ELECTRA 的 RTD:
- MLM 只预测被掩盖的 token
- RTD(Replaced Token Detection)判别每个 token 是否被替换
- RTD 通常计算效率更高
提升 MLM 效果的架构调整
- 使用动态掩码:每次 epoch 重新随机掩盖
- 调整掩盖比例(通常 15% 最佳)
- 部分掩盖词使用原词,部分使用随机词
- 使用全词掩码(Whole Word Masking)
开放性问题
- 如何设计实验验证 MLM 确实学习到了语言理解能力,而不是简单的模式匹配?
- 对于低资源语言,哪些 MLM 的改进策略可能特别有效?
- MLM 预训练和下游任务 fine-tuning 之间是否存在知识冲突?如何缓解?
总结
本文详细解析了 BERT 中 MLM 任务的核心公式和实现细节。通过理解 softmax 交叉熵损失函数、掩码机制以及 PyTorch 实现代码,读者可以掌握这一重要技术。实践部分的学习率设置、长文本处理技巧和常见问题解决方案,为实际应用提供了实用指导。最后的延伸思考和开放性问题,希望能激发读者对 MLM 更深入的探索。
MLM 作为 BERT 的核心预训练任务,其设计理念和实现细节值得深入研究和理解。随着对这部分内容的掌握,读者可以更好地应用和调整 BERT 模型,甚至设计新的预训练任务,推动 NLP 技术的发展。
