BERT掩码语言模型实战入门:从原理到PyTorch实现

1次阅读
没有评论

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

image.webp

为什么需要掩码语言模型?

在自然语言处理(NLP)领域,传统的 n -gram 语言模型存在两个致命缺陷:

BERT 掩码语言模型实战入门:从原理到 PyTorch 实现

  1. 上下文窗口固定:n-gram 只能看到前后固定数量的词(比如 3 -gram),无法建模长距离依赖关系
  2. 数据稀疏问题:随着 n 增大,组合爆炸导致很多序列在训练集中从未出现过

BERT 的掩码语言模型(Masked Language Model, MLM)通过随机遮盖输入文本中的部分单词(如 15%),让模型根据双向上下文预测被遮盖的词。这种自监督训练方式让模型学会深度理解词语在上下文中的真实含义。

核心实现四部曲

1. 输入数据预处理

关键操作流程:

  1. 使用 WordPiece 分词器将文本转为 token ID 序列
  2. 随机选择 15% 的 token 进行替换,其中:
  3. 80% 替换为[MASK]
  4. 10% 替换为随机 token
  5. 10% 保持原 token(增加模型纠错能力)
  6. 添加 [CLS] 和[SEP]等特殊 token
def mask_tokens(inputs, tokenizer, mask_prob=0.15):
    """动态生成掩码样本"""
    labels = inputs.clone()
    # 生成掩码位置矩阵(1 表示需要 mask)prob_matrix = torch.full(labels.shape, mask_prob)
    masked_indices = torch.bernoulli(prob_matrix).bool()

    # 80% 概率替换为[MASK]
    indices_replaced = torch.bernoulli(torch.full(labels.shape, 0.8)).bool() & masked_indices
    inputs[indices_replaced] = tokenizer.mask_token_id

    # 10% 概率替换为随机词
    indices_random = torch.bernoulli(torch.full(labels.shape, 0.5)).bool() & masked_indices & ~indices_replaced
    random_words = torch.randint(len(tokenizer), labels.shape, dtype=torch.long)
    inputs[indices_random] = random_words[indices_random]

    return inputs, labels

2. 注意力掩码构建

BERT 需要处理变长输入,通过 attention_mask 区分有效内容与 padding 部分:

  • 1 表示真实 token
  • 0 表示 padding 部分
# 假设输入序列长度为 128,实际有效长度 100
token_ids = [101, 2053, ..., 102, 0, 0, ..., 0]  # 含 28 个 padding
attention_mask = [1]*100 + [0]*28  # 前 100 位为 1,padding 为 0 

3. 模型前向传播

注意维度变换的关键点:

from transformers import BertForMaskedLM

model = BertForMaskedLM.from_pretrained('bert-base-uncased')

# 输入尺寸: [batch_size, seq_len]
outputs = model(
    input_ids,
    attention_mask=attention_mask,
    labels=labels  # 可选,计算 loss 时需传入
)
# logits 尺寸: [batch_size, seq_len, vocab_size]
logits = outputs.logits

4. 损失计算优化

原始 BERT 的 loss 只计算被 mask 位置的预测结果:

loss_fct = torch.nn.CrossEntropyLoss()
# 只计算 mask 位置的 loss(忽略 padding 和未 mask 部分)active_loss = attention_mask.view(-1) == 1  # 有效 token 位置
active_logits = logits.view(-1, model.config.vocab_size)[active_loss]
active_labels = labels.view(-1)[active_loss]
loss = loss_fct(active_logits, active_labels)

性能优化实战技巧

掩码比例影响

通过实验发现不同 mask 概率的影响:

  • 15%:原始论文推荐值,预测难度适中
  • 25%:提升任务难度,可能需要更长时间训练
  • 30%:信息丢失严重,可能损害模型性能

混合精度训练

使用 NVIDIA 的 Apex 库实现 FP16 训练:

  1. 初始化缩放器
    from apex import amp
    model, optimizer = amp.initialize(model, optimizer, opt_level="O1")
  2. 修改 loss 计算
    with amp.scale_loss(loss, optimizer) as scaled_loss:
        scaled_loss.backward()
  3. 梯度裁剪仍需在缩放前进行
    torch.nn.utils.clip_grad_norm_(amp.master_params(optimizer), max_grad_norm)

延伸思考

改进静态掩码策略

原始 BERT 在预处理时一次性生成所有 mask,可能导致:
– 相同数据多次训练时看到相同的 mask 模式
– 解决方案:动态掩码(如本文实现),每个 epoch 重新生成

特殊 token 处理

[CLS]和 [SEP] 等特殊 token 通常不参与 mask,因为:
1. [CLS]需要聚合全局信息用于分类任务
2. [SEP]标记句子边界,破坏结构影响性能
3. 可通过修改 mask 生成逻辑排除这些位置

# 排除特殊 token 的 mask 示例
special_tokens = [tokenizer.cls_token_id, tokenizer.sep_token_id]
maskable_positions = ~torch.isin(inputs, torch.tensor(special_tokens))
prob_matrix = prob_matrix * maskable_positions.float()

结语

通过这次实践,我深刻体会到 BERT 的掩码策略设计之精妙——既不像传统语言模型那样 ” 偷看 ” 答案,又能通过双向上下文学习词语的深层语义。建议初学者可以先用小批量数据(如 5000 条)跑通整个流程,再逐步扩展到大数据集训练。

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