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

微调BertForPreTraining时出现CUDA OOM显存不足错误如何解决?

BertForPreTraining微调显存溢出问题解决方案

核心原因

bert-base-uncased单模型FP32精度下显存占用约1.3G,训练过程中还需要存储梯度、优化器状态、batch计算中间张量,哪怕batch size=2、序列长度512的场景,最低显存要求也在6G以上,2G/4G显存无法满足默认配置的训练需求。

可落地修改方案

  • 启用混合精度训练
    引入PyTorch自带的AMP模块,可降低近50%显存占用,修改代码如下:
    在优化器初始化后添加scaler定义:
    scaler = torch.cuda.amp.GradScaler()
    
    训练循环内前向、反向传播部分修改为:
    with torch.cuda.amp.autocast():
        outputs = model(input_ids, attention_mask=attention_mask,
                        token_type_ids=token_type_ids,
                        next_sentence_label=next_sentence_label,
                        labels=labels)
    loss = outputs.loss
    scaler.scale(loss).backward()
    scaler.step(optim)
    scaler.update()
    
  • 降低输入序列长度
    当前max_length=512下自注意力计算的显存占用和序列长度成平方关系,若你的语料句子偏短,可将tokenizer调用时的max_length参数调整为128或256,显存占用可下降75%~50%。
  • 启用梯度检查点
    模型移到GPU前添加一行代码开启梯度检查点,可节省30%左右显存,仅会带来少量训练速度损失:
    model.gradient_checkpointing_enable()
    
  • 更换轻量预训练模型
    可替换bert-base-uncased为distilbert-base-uncased,模型体积减小40%,训练速度提升60%,预训练任务性能仅下降约3%,完全可以在4G显存下正常训练。
  • 梯度累加
    若需要维持等效大batch size,可加入梯度累加逻辑,示例如下:
    accum_steps = 4 # 等效batch size=2*4=8
    for step, batch in enumerate(loop):
        # 前向反向逻辑不变
        loss = outputs.loss / accum_steps
        loss.backward()
        if (step+1) % accum_steps == 0:
            optim.step()
            optim.zero_grad()
    
  • 修复代码警告
    Dataset类中的torch.tensor(val[idx])写法会带来不必要的张量复制,修改为:
    return {key: val[idx].clone().detach() for key, val in self.encodings.items()}
    

内容的提问来源于stack exchange,提问作者Rastin Bamdad

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.27 07:06:03