深入解析BERT掩码语言模型(MLM)目标函数:从理论到PyTorch实现

1次阅读
没有评论

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

image.webp

BERT 掩码语言模型 (MLM) 目标函数全解析

背景与核心思想

BERT(Bidirectional Encoder Representations from Transformers)的核心创新之一就是采用了掩码语言模型(Masked Language Model, MLM)作为预训练目标。与传统的单向语言模型不同,MLM 通过随机掩盖输入序列中的部分 token(标记),让模型基于上下文双向预测被掩盖的内容。这种设计使 BERT 能够捕获更丰富的上下文信息。

深入解析 BERT 掩码语言模型 (MLM) 目标函数:从理论到 PyTorch 实现

数学原理

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)

实现差异分析

  1. 原始论文实现
  2. 严格遵循 15% 的掩码比例
  3. 80-10-10 的替换策略(80%[MASK], 10% 随机 token, 10% 保持原词)

  4. HuggingFace 实现

  5. 支持动态掩码(每次 epoch 重新生成掩码模式)
  6. 更灵活的特殊 token 处理
  7. 集成了子词 (token) 边界处理逻辑

调优技巧

  1. 掩码比例
  2. 15% 是经验值,可在 10%-20% 间调整
  3. 对于专业领域文本,可适当降低比例

  4. 罕见词处理

  5. 对于低频词,可提高其被掩码的概率
  6. 使用子词(subword)tokenizer 减少 OOV 问题

  7. 混合精度训练

  8. 使用 torch.cuda.amp 自动管理精度
  9. 注意 log_softmax 的数值稳定性

避坑指南

  1. 常见错误
  2. 错误计算有效 token 数(未忽略 padding 位置)
  3. 未正确处理特殊 token 导致的训练偏差

  4. 多 GPU 训练

  5. 确保各 GPU 的掩码模式不同
  6. 梯度同步时注意 batch norm 统计量

  7. 验证指标

  8. MLM 准确率与下游任务性能非严格正相关
  9. 建议同时监控困惑度(perplexity)

延伸思考

  1. 如何设计动态掩码策略,使模型在不同训练阶段看到不同难度的样本?
  2. 对于中文等连续文本语言,是否应该调整掩码单位(如整词掩码)?
  3. 如何结合 MLM 和 ELECTRA 的替换 token 检测 (RTD) 目标?

参考文献

  1. BERT 原始论文:arXiv:1810.04805
  2. HuggingFace Transformers 实现
  3. ELECTRA 论文:arXiv:2003.10555
正文完
 0
评论(没有评论)