深入解析BERT掩码语言模型(MLM)核心公式:从数学原理到工程实现

1次阅读
没有评论

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

image.webp

背景介绍

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

深入解析 BERT 掩码语言模型 (MLM) 核心公式:从数学原理到工程实现

核心公式解析

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 最有效)

避坑指南

常见实现错误

  1. mask 比例不当:通常 15% 的 mask 比例效果最佳,过高会导致模型难以学习,过低则训练不充分
  2. 静态 mask:应在每个 epoch 重新随机 mask,避免模型记忆特定模式
  3. 忽略特殊 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 范围

开放性问题

  1. 如何设计更有效的 mask 策略(如 span masking)来提升模型性能?
  2. 在大规模预训练中,如何平衡计算效率和模型表现?
  3. MLM 损失与其他辅助损失(如 NSP)如何协同优化?
正文完
 0
评论(没有评论)