BERT多标签不平衡数据集训练出现NaN损失的原因及解决方法咨询
基于BERT的不平衡多标签文本分类训练NaN损失问题求助
数据集配置
- 样本量约100K
- 标签数约50
- 不平衡情况:部分标签覆盖80%样本,部分标签占比不足0.5%
模型配置
使用Hugging Face的bert-base-uncased搭配自定义分类头,代码如下:
from transformers import BertModel import torch.nn as nn class MultiLabelClassifier(nn.Module): def __init__(self, num_labels): super(MultiLabelClassifier, self).__init__() self.bert = BertModel.from_pretrained("bert-base-uncased") self.classifier = nn.Linear(self.bert.config.hidden_size, num_labels) def forward(self, input_ids, attention_mask): outputs = self.bert(input_ids=input_ids, attention_mask=attention_mask) logits = self.classifier(outputs.pooler_output) # 使用[CLS] token return logits
为处理数据不平衡,采用带权重的BCEWithLogitsLoss:
class_weights = torch.tensor([1.0 / (freq + 1e-5) for freq in label_frequencies]).to(device) criterion = nn.BCEWithLogitsLoss(pos_weight=class_weights)
问题表现
- 初始损失正常,训练2-3个epoch后损失突然变为NaN
- NaN出现前部分logits值异常(如1e20、-1e20)
- NaN出现前伴随梯度爆炸现象
已尝试的解决方法
- 添加梯度裁剪(
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)),仅略有改善 - 将学习率降至5e-6,仅延迟NaN出现时间,未彻底解决
- 重新初始化分类器权重:
for layer in model.classifier.parameters(): if isinstance(layer, nn.Linear): nn.init.xavier_normal_(layer.weight)
- 用torch.logsigmoid重写损失计算以提升数值稳定性:
loss = -labels * torch.logsigmoid(logits) - (1 - labels) * torch.logsigmoid(-logits) loss = loss.mean()
以上方法均未能彻底解决问题。
疑问
- 问题成因是什么?是否由数据集极端不平衡导致?
- 如何彻底解决?是否可尝试标签平滑,或有更优的训练稳定方案?
训练循环片段(参考)
optimizer = AdamW(model.parameters(), lr=5e-5) scheduler = torch.optim.lr_scheduler.LinearLR(optimizer, total_iters=num_training_steps) for epoch in range(num_epochs): model.train() for batch in dataloader: input_ids, attention_mask, labels = batch["input_ids"], batch["attention_mask"], batch["labels"] logits = model(input_ids, attention_mask) loss = criterion(logits, labels) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() scheduler.step()
需求
寻求可彻底避免NaN损失、稳定不平衡多标签数据集训练的实用方案,恳请有相关经验的开发者提供帮助。
内容的提问来源于stack exchange,提问作者Erhan Arslan
相关产品推荐
相关产品推荐

