微调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
相关产品推荐
相关产品推荐

