两次调用模型引发内存错误,求不破坏计算图的解决方案
解决方案
针对两次调用模型导致的内存错误,同时需要保留两个loss的梯度计算,可尝试以下几种方案:
1. 拆分模型共享层与分支层,避免重复计算
如果mode='a'和mode='b'的前向流程存在共享的网络层,将这部分单独抽离,先计算一次共享层输出,再分别走两个分支,能大幅减少重复计算带来的内存占用:
# 重构模型,拆分共享部分和分支部分 shared_out = model.shared_layers(input) output1 = model.branch_a(shared_out) loss1 = loss_fun1(output1) output2 = model.branch_b(shared_out) loss2 = loss_fun2(output2) total_loss = loss1 + loss2 optimizer.zero_grad() total_loss.backward() optimizer.step()
2. 使用梯度检查点(Gradient Checkpointing)
通过torch.utils.checkpoint.checkpoint可以在正向传播时不存储所有中间激活张量,而是在反向传播时重新计算这些激活,以此节省内存。需要将模型的不同mode分支封装成可调用函数:
from torch.utils.checkpoint import checkpoint def forward_a(model, input): return model(input, mode='a') def forward_b(model, input): return model(input, mode='b') # 使用checkpoint包裹前向传播 output1 = checkpoint(forward_a, model, input) loss1 = loss_fun1(output1) output2 = checkpoint(forward_b, model, input) loss2 = loss_fun2(output2) total_loss = loss1 + loss2 optimizer.zero_grad() total_loss.backward() optimizer.step()
注意:梯度检查点会增加一点反向传播的计算时间,属于内存与速度的 trade-off。
3. 手动清理无用张量并释放缓存
在计算完第一个分支后,手动删除不再需要的中间张量,并调用缓存释放接口,及时回收内存:
output1 = model(input, mode='a') loss1 = loss_fun1(output1) # 删除output1,释放其占用的内存 del output1 # 如果使用GPU,清空CUDA缓存 torch.cuda.empty_cache() output2 = model(input, mode='b') loss2 = loss_fun2(output2) total_loss = loss1 + loss2 optimizer.zero_grad() total_loss.backward() optimizer.step()
此方法简单直接,但效果取决于模型中间张量的内存占比,适合临时应急。
4. 启用混合精度训练
使用PyTorch的自动混合精度(AMP)功能,将部分张量以半精度(FP16)存储,在不影响梯度计算精度的前提下降低内存消耗:
from torch.cuda.amp import GradScaler, autocast scaler = GradScaler() optimizer.zero_grad() with autocast(): output1 = model(input, mode='a') loss1 = loss_fun1(output1) output2 = model(input, mode='b') loss2 = loss_fun2(output2) total_loss = loss1 + loss2 # 反向传播与优化 scaler.scale(total_loss).backward() scaler.step(optimizer) scaler.update()
该方案需要GPU支持FP16计算,现在主流NVIDIA显卡都兼容。
内容的提问来源于stack exchange,提问作者user491683
相关产品推荐
相关产品推荐

