文章目录

BERT的Masked Language Model损失函数与NSP任务

发布于 2026-07-07 12:43:03 · 浏览 66 次 · 评论 0 条

BERT的Masked Language Model损失函数与NSP任务

理解BERT的核心在于弄懂它预训练时的两个核心任务及其对应的损失函数。这相当于弄清楚模型在学习过程中如何“做题”和“对答案”。


1. 理解并实现 Masked Language Model (MLM) 任务

MLM是BERT的灵魂任务,其目标是让模型根据上下文预测被随机掩盖的词。你可以把它想象成一个高级的完形填空游戏。

  1. 构造 输入数据。
    从文本语料库中获取一个句子或句子对。随机选择句子中 15% 的词(Token)作为要预测的目标。对于这 15% 的选中词,执行以下操作:

    • 80% 的概率:替换为特殊的 [MASK] 标记。
    • 10% 的概率:替换为另一个随机的词。
    • 10% 的概率:保持原词不变。
  2. 输入 处理后的序列到BERT模型。
    整个句子(包含可能被替换的词)被输入模型。模型的输出是每个位置对应整个词汇表的概率分布。

  3. 计算 MLM损失。
    损失函数只计算那 15% 被选中位置的预测误差。对于每一个被选中的位置,模型输出一个概率分布 P,其中正确词 w 对应的概率为 P_w。MLM任务的损失是这些位置上预测概率的负对数似然之和。
    其核心数学表达为:

    $$ \mathcal{L}_{\text{MLM}} = -\sum_{i \in \mathcal{M}} \log P(w_i | \mathbf{w}_{\backslash i}) $$

    这个公式的大白话解释是:对于所有被掩盖的位置 i(属于集合 \mathcal{M}),损失函数会惩罚模型给出的错误预测的概率。P(w_i | \mathbf{w}_{\backslash i}) 代表在已知所有其他词\mathbf{w}_{\backslash i})的条件下,模型预测出正确词 w_i 的概率。我们希望这个概率尽可能高,因此通过取对数并取负,将最大化概率的问题转化为了最小化损失的问题。

  4. 优化 模型参数。
    通过反向传播算法,根据计算出的 \mathcal{L}_{\text{MLM}}调整模型内部数以亿计的参数,使得下一次面对类似填空题时,模型能给出更准的答案。


2. 理解并实现 Next Sentence Prediction (NSP) 任务

NSP是BERT的另一个辅助任务,旨在让模型理解句子之间的逻辑关系。它训练模型判断两个句子是否在原文中是相邻关系。

  1. 准备 句子对数据。
    从语料库中构造两种类型的样本,每种各占 50%

    • 正样本(IsNext):句子B确实是句子A在原文中的下一句。
    • 负样本(NotNext):句子B是从语料库中随机抽取的一个句子,与句子A无关。
  2. 输入 句子对到BERT模型。
    将句子A和句子B用特殊的 [SEP] 标记连接起来,输入模型。模型在 [CLS] 标记位置的输出向量被用作整个句子对的表征。

  3. 计算 NSP损失。
    [CLS] 表征之上,添加一个简单的分类层,用于输出一个概率值,表示句子B是句子A的下一个句子的可能性。这是一个标准的二元分类问题,其损失函数是二元交叉熵损失。
    其数学表达为:

    $$ \mathcal{L}_{\text{NSP}} = -[y \log p + (1 - y) \log (1 - p)] $$

    其中,y 是标签(如果是真实下一句,则 y=1;如果是随机句子,则 y=0)。p 是模型预测的“是下一句”的概率。这个损失函数会推动模型对于正样本给出更高的 p,对于负样本给出更低的 p

  4. 综合 两个任务进行联合训练。
    在实际预训练中,BERT的总损失是两个任务损失的加权和,通常是直接相加。模型同时优化这两个目标,从而既能精准理解词与上下文的关系(来自MLM),也能初步把握句子间的连贯性(来自NSP)。

评论 (0)

暂无评论,快来抢沙发吧!

扫一扫,手机查看

扫描上方二维码,在手机上查看本文