预训练BERT Token分类任务训练时内存溢出问题求助
嵌套NER模型Ubuntu训练内存溢出问题排查与修复
问题背景
基于预训练BertForTokenClassification实现嵌套命名实体识别模型,采用多标签方法输出3组层级logits。在Ubuntu 22.04(16GB内存)训练时,无论batch_size设为1还是16,内存占用都会持续增长,最终触发OutOfMemory错误;但相同内存配置的Windows 11环境下无此问题,仅训练速度慢3倍。
核心修复方案
1. 修复训练损失累积导致的计算图泄漏
训练循环中直接累加带梯度的total_loss张量,会持续累积计算图,导致内存占用不断增长。
修改训练模块的train方法:
# 原代码 train_loss += total_loss # ... train_loss = train_loss / num_training_steps # 修改为 train_loss += total_loss.item() # 累加数值而非张量,避免保留计算图 # ... train_loss = train_loss / len(train_loader) # 修正平均损失计算逻辑,除以batch数而非总步数
2. 删除评估函数中冗余的标签存储
评估函数中true_labels和pred_labels被持续填充但未被使用,累积大量冗余数据占用内存。
简化后的evaluate方法:
def evaluate(self, model, dataloader): eval_loss = 0 for input_ids, attention_mask, labels in dataloader: input_ids, attention_mask = input_ids.to(self.device), attention_mask.to(self.device) labels_pred = model(input_ids, attention_mask) total_loss = 0 for i in range(3): loss = F.cross_entropy(labels_pred[i].view(-1, 8), labels[:, i].reshape(-1), ignore_index=0) total_loss += loss eval_loss += total_loss.item() # 改为数值累加 eval_loss /= len(dataloader) return eval_loss
3. 移除模型初始化的冗余参数
当前使用BertForTokenClassification加载预训练权重,该类自带的分类头未被使用,导致冗余参数占用内存。
修改模型类的初始化代码:
# 原代码 from transformers import BertForTokenClassification self.bert = BertForTokenClassification.from_pretrained(model_name, hidden_dropout_prob=dropout) # 修改为 from transformers import BertModel self.bert = BertModel.from_pretrained(model_name, hidden_dropout_prob=dropout)
额外优化建议
- 数据加载器调整:若DataLoader设置了
num_workers>0,尝试改为num_workers=0(主进程加载数据),避免Linux多进程内存泄漏; - 版本升级:对齐Ubuntu与Windows环境的PyTorch、Transformers版本,升级至稳定版:
pip install --upgrade torch transformers
内容的提问来源于stack exchange,提问作者Alexandr Duck
相关产品推荐
相关产品推荐

