You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

预训练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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.16 21:04:53