共计 2031 个字符,预计需要花费 6 分钟才能阅读完成。
1. BERT 模型损失函数基础解析
BERT 模型的预训练阶段依赖两个核心任务:掩码语言模型(MLM)和下一句预测(NSP),其损失函数设计直接影响模型的语义表征能力。

1.1 MLM 任务损失函数
给定输入序列 $X={x_1,…,x_n}$,随机掩码 15% 的 token 得到 $X^{mask}$,模型需预测被掩码位置的原始 token。其损失函数为交叉熵损失:
$$\mathcal{L}{MLM} = -\sum)$$} \log P(x_i|X^{mask
其中 $M$ 表示被掩码的 token 位置集合。实践中采用以下改进策略:
- 80% 概率替换为 [MASK] 标记
- 10% 概率替换为随机 token
- 10% 概率保持原 token 不变
1.2 NSP 任务损失函数
对于句子对 $(A,B)$,模型需判断 B 是否为 A 的下一句,损失函数为二分类交叉熵:
$$\mathcal{L}_{NSP} = -[y\log p + (1-y)\log(1-p)]$$
其中 $y\in{0,1}$ 表示是否为连续句子,$p$ 为模型预测概率。后续研究发现 NSP 任务对某些下游任务帮助有限,RoBERTa 等模型已移除此任务。
2. 损失函数变体对比
2.1 Label Smoothing
传统交叉熵使用硬标签(0 或 1),容易导致过拟合。标签平滑通过引入平滑因子 $\epsilon$ 软化标签:
$$q_i = \begin{cases}
1-\epsilon + \epsilon/K & \text{if} i=y \
\epsilon/K & \text{otherwise}
\end{cases}$$
其中 $K$ 为类别数,$\epsilon$ 通常取 0.1。在 GLUE 基准测试中,该方法可使模型平均提升 0.3-0.5 个点。
2.2 Focal Loss
针对类别不平衡问题,Focal Loss 通过调制因子 $(1-p_t)^\gamma$ 降低易分类样本的权重:
$$FL(p_t) = -\alpha_t(1-p_t)^\gamma \log(p_t)$$
当 $\gamma=2$ 时,在低资源文本分类任务中可提升 2 -5% 的 F1 值。
3. PyTorch 实现示例
3.1 标准交叉熵实现
import torch.nn as nn
class BertMLMLoss(nn.Module):
def __init__(self, vocab_size):
super().__init__()
self.loss_fn = nn.CrossEntropyLoss(ignore_index=-100) # 忽略 padding 位置
def forward(self, logits, labels):
# logits: [batch_size, seq_len, vocab_size]
# labels: [batch_size, seq_len]
loss = self.loss_fn(logits.view(-1, logits.size(-1)),
labels.view(-1)
)
return loss
3.2 带权重损失变体
class WeightedMLMLoss(nn.Module):
def __init__(self, vocab_size, class_weights):
super().__init__()
self.weights = torch.FloatTensor(class_weights)
self.loss_fn = nn.CrossEntropyLoss(weight=self.weights)
def forward(self, logits, labels):
# 对低频词赋予更高权重
return self.loss_fn(logits.view(-1, logits.size(-1)), labels.view(-1))
3.3 超参数调优建议
- 学习率与损失缩放:使用 AdamW 优化器时,建议初始 lr=2e-5
- 梯度裁剪:设置 max_grad_norm=1.0 防止梯度爆炸
- 混合精度训练:搭配 AMP 自动损失缩放可提升 20% 训练速度
4. 生产环境最佳实践
4.1 损失监控方案
- 实时监控各任务损失曲线
- 设置损失阈值告警(如连续 3 个 epoch 增加 >5%)
- 使用 WandB/TensorBoard 可视化
4.2 多任务损失平衡
采用动态加权策略:
$$w_k(t) = \frac{\exp(-\alpha \cdot \bar{L}_k)}{\sum_i \exp(-\alpha \cdot \bar{L}_i)}$$
其中 $\bar{L}_k$ 为任务 k 的滑动平均损失,$\alpha$ 控制调整强度。
4.3 分布式训练同步
- 使用 torch.nn.parallel.DistributedDataParallel
- 确保所有进程的损失计算一致
- 梯度聚合前进行 all_reduce 操作
5. 开放性问题讨论
当处理类别极度不平衡的文本分类(如欺诈检测)时:
1. 如何设计基于样本难度的自适应损失函数?
2. 能否将对比学习损失融入 BERT 微调阶段?
3. 预训练任务损失与下游任务损失的最优比例如何确定?
期待读者在实践中探索这些问题的解决方案。
