共计 2672 个字符,预计需要花费 7 分钟才能阅读完成。
背景痛点
在 NLP 任务中,BERT 模型的微调阶段常遇到两个典型问题:

-
梯度消失:深层 Transformer 架构中,传统的交叉熵损失容易导致底层参数更新缓慢。在文本分类任务中,当类别分布极度不平衡时(如负面评论占比 5%),模型可能过早陷入局部最优。
-
标签噪声敏感:人工标注的文本数据常存在约 3 -5% 的错误标签。标准交叉熵会强制模型拟合这些噪声样本,反而降低泛化能力。我们的实验显示,在 AG News 数据集上,10% 的随机标签噪声可使测试准确率下降 8%。
技术对比
三类损失函数特性对比
| 损失类型 | 优点 | 缺点 | 适用场景 |
|---|---|---|---|
| 标准交叉熵 | 计算高效,理论完备 | 对噪声敏感 | 类别均衡的干净数据 |
| Focal Loss | 缓解类别不平衡 | 需手动调节 γ 参数 | 长尾分布数据集 |
| 对比损失 | 增强类间区分度 | 计算复杂度较高 | 细粒度分类任务 |
数学表达对比:
– 标准交叉熵:$L_{CE} = -\sum_{i=1}^C y_i \log(p_i)$
– Focal Loss:$L_{FL} = -\sum_{i=1}^C (1-p_i)^\gamma y_i \log(p_i)$
– 对比损失:$L_{con} = -\log \frac{e^{s_p/\tau}}{e^{s_p/\tau} + \sum_{n=1}^N e^{s_n/\tau}}}$
核心实现
带标签平滑的交叉熵
import torch
import torch.nn as nn
import torch.nn.functional as F
class LabelSmoothingCE(nn.Module):
def __init__(self, smoothing=0.1, dim=-1):
super().__init__()
self.smoothing = smoothing
self.dim = dim
def forward(self, logits: torch.Tensor, targets: torch.Tensor):
assert 0 <= self.smoothing < 1
with torch.no_grad():
# 形状检查
if targets.dim() != logits.dim() - 1:
raise ValueError(f"Target dim {targets.dim()} != logits dim-1 {logits.dim()-1}")
# 构建平滑标签
num_classes = logits.size(self.dim)
true_dist = torch.full_like(logits, self.smoothing/(num_classes-1))
true_dist.scatter_(1, targets.unsqueeze(1), 1-self.smoothing)
return torch.mean(torch.sum(-true_dist * F.log_softmax(logits, dim=self.dim), dim=self.dim))
温度系数可调的对比损失
class ContrastiveLoss(nn.Module):
def __init__(self, temp=0.5, eps=1e-8):
super().__init__()
self.temp = temp
self.eps = eps
def forward(self, features: torch.Tensor, labels: torch.Tensor):
device = features.device
batch_size = features.shape[0]
# 归一化特征向量
features = F.normalize(features, p=2, dim=1)
# 计算相似度矩阵
sim_matrix = torch.matmul(features, features.T) / self.temp
# 构建正负样本掩码
labels = labels.contiguous().view(-1,1)
mask = torch.eq(labels, labels.T).to(device)
pos_mask = mask.fill_diagonal_(False) # 排除自身
# 计算对比损失
exp_sim = torch.exp(sim_matrix)
pos_sum = torch.sum(exp_sim * pos_mask, dim=1)
neg_sum = torch.sum(exp_sim * (1-mask), dim=1)
loss = -torch.log(pos_sum/(pos_sum + neg_sum + self.eps)).mean()
return loss
实验验证
在 GLUE 的 MRPC 数据集上的实验结果:
| 损失类型 | 训练时间(min) | 显存占用(GB) | 准确率(%) |
|---|---|---|---|
| 标准交叉熵 | 42 | 3.8 | 86.2 |
| 标签平滑(α=0.1) | 38 | 3.8 | 87.1 |
| 对比损失(τ=0.3) | 51 | 4.2 | 87.9 |
关键发现:
1. 标签平滑使训练速度提升 9%,主要得益于更稳定的梯度
2. 对比损失虽然耗时增加,但在 F1 分数上比基线高 1.7 个点
避坑指南
- 梯度裁剪:
- BERT 推荐阈值 3.0,过大失去约束效果,过小会抑制学习
-
监控梯度范数:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=3.0) -
混合精度训练:
- 使用
torch.cuda.amp.GradScaler()自动处理 loss scaling -
对 softmax 计算需保持 fp32 精度:
with autocast(dtype=torch.float16): # 前向计算 logits = model(input_ids) # 手动转换损失计算精度 loss = criterion(logits.float(), labels) -
分布式训练:
- 对比损失需要同步所有 GPU 的特征向量:
# 使用 DistributedDataParallel 时 features_list = [torch.zeros_like(features) for _ in range(world_size)] torch.distributed.all_gather(features_list, features) all_features = torch.cat(features_list, dim=0)
优化建议
- 动态温度系数:初期用较大 τ(如 1.0)探索空间,后期逐渐降低到 0.3
- 损失组合:尝试交叉熵 + 对比损失的加权组合,比例建议 4:1
- 对于超长文本,在计算对比损失时采用随机采样策略降低计算量
通过以上方法,我们在客户服务工单分类任务中实现了 23% 的训练加速,同时保持了分类准确率。关键在于根据数据特性选择合适的损失函数组合,并做好训练过程的监控与调优。
