共计 2452 个字符,预计需要花费 7 分钟才能阅读完成。
BERT 掩码语言模型 (MLM) 目标函数全解析
背景与核心思想
BERT(Bidirectional Encoder Representations from Transformers)的核心创新之一就是采用了掩码语言模型(Masked Language Model, MLM)作为预训练目标。与传统的单向语言模型不同,MLM 通过随机掩盖输入序列中的部分 token(标记),让模型基于上下文双向预测被掩盖的内容。这种设计使 BERT 能够捕获更丰富的上下文信息。

数学原理
MLM 的损失函数采用标准的交叉熵(Cross-Entropy)形式。对于被掩盖的 token,其损失计算如下:
$$
\mathcal{L}{MLM} = -\sum)
$$} \log P(x_i | x_{\backslash i
其中:
– $masked$ 表示被掩盖的 token 位置集合
– $x_i$ 是第 i 个位置的真实 token
– $x_{\backslash i}$ 表示除第 i 个位置外的所有输入 token
PyTorch 实现细节
1. 输入序列的随机掩码处理
import torch
import random
def random_mask(input_ids, mask_token_id, vocab_size, mask_prob=0.15):
"""
对输入序列进行随机掩码处理
:param input_ids: 输入 token ID 序列 [batch_size, seq_len]
:param mask_token_id: [MASK]标记的 ID
:param vocab_size: 词表大小
:param mask_prob: 掩码概率
:return: 处理后的 input_ids, 掩码位置的 labels
"""
labels = input_ids.clone()
# 创建随机掩码矩阵
probability_matrix = torch.full(labels.shape, mask_prob)
# 特殊 token 不参与掩码(如 [CLS], [SEP] 等)special_tokens_mask = (input_ids == 0) | (input_ids == 1) # 示例,实际需根据 tokenizer 调整
probability_matrix.masked_fill_(special_tokens_mask, value=0.0)
masked_indices = torch.bernoulli(probability_matrix).bool()
labels[~masked_indices] = -100 # 只计算被掩盖位置的损失
# 80% 的概率替换为[MASK]
indices_replaced = torch.bernoulli(torch.full(labels.shape, 0.8)).bool() & masked_indices
input_ids[indices_replaced] = mask_token_id
# 10% 的概率替换为随机 token
indices_random = torch.bernoulli(torch.full(labels.shape, 0.5)).bool() & masked_indices & ~indices_replaced
random_words = torch.randint(vocab_size, labels.shape, dtype=torch.long)
input_ids[indices_random] = random_words[indices_random]
return input_ids, labels
2. 损失计算实现
import torch.nn as nn
def mlm_loss(logits, labels):
"""
计算 MLM 损失
:param logits: 模型输出 [batch_size, seq_len, vocab_size]
:param labels: 标签 [batch_size, seq_len],其中非掩码位置为 -100
:return: 标量损失值
"""
loss_fct = nn.CrossEntropyLoss(ignore_index=-100)
# 将 logits 和 labels 展平
active_loss = labels.view(-1) != -100
active_logits = logits.view(-1, logits.size(-1))[active_loss]
active_labels = labels.view(-1)[active_loss]
return loss_fct(active_logits, active_labels)
实现差异分析
- 原始论文实现:
- 严格遵循 15% 的掩码比例
-
80-10-10 的替换策略(80%[MASK], 10% 随机 token, 10% 保持原词)
-
HuggingFace 实现:
- 支持动态掩码(每次 epoch 重新生成掩码模式)
- 更灵活的特殊 token 处理
- 集成了子词 (token) 边界处理逻辑
调优技巧
- 掩码比例:
- 15% 是经验值,可在 10%-20% 间调整
-
对于专业领域文本,可适当降低比例
-
罕见词处理:
- 对于低频词,可提高其被掩码的概率
-
使用子词(subword)tokenizer 减少 OOV 问题
-
混合精度训练:
- 使用 torch.cuda.amp 自动管理精度
- 注意 log_softmax 的数值稳定性
避坑指南
- 常见错误:
- 错误计算有效 token 数(未忽略 padding 位置)
-
未正确处理特殊 token 导致的训练偏差
-
多 GPU 训练:
- 确保各 GPU 的掩码模式不同
-
梯度同步时注意 batch norm 统计量
-
验证指标:
- MLM 准确率与下游任务性能非严格正相关
- 建议同时监控困惑度(perplexity)
延伸思考
- 如何设计动态掩码策略,使模型在不同训练阶段看到不同难度的样本?
- 对于中文等连续文本语言,是否应该调整掩码单位(如整词掩码)?
- 如何结合 MLM 和 ELECTRA 的替换 token 检测 (RTD) 目标?
参考文献
- BERT 原始论文:arXiv:1810.04805
- HuggingFace Transformers 实现
- ELECTRA 论文:arXiv:2003.10555
正文完
