如何解决PyTorch保存LLM微调Checkpoint时的CUDA OOM异常?
解决PyTorch保存LLM Checkpoint时的OOM问题
完全可以通过将模型参数转移到CPU内存后再保存Checkpoint来解决这个GPU显存不足的问题。保存Checkpoint时PyTorch默认会在GPU端处理数据,额外占用显存空间,把参数转移到CPU后,只会占用系统内存,不会消耗GPU显存。
以下是两种实用的实现方法:
手动迁移模型到CPU后保存
先将整个模型从GPU转移到CPU,再执行保存操作,示例代码:# 把模型从GPU移到CPU model.to('cpu') # 保存模型参数 torch.save(model.state_dict(), 'llm_checkpoint.pth') # 若后续还要继续训练,记得把模型移回GPU model.to('cuda')保存时指定
map_location参数
无需手动移动模型,直接在torch.save()中通过map_location参数将参数映射到CPU,这种方式更简洁,示例代码:torch.save(model.state_dict(), 'llm_checkpoint.pth', map_location='cpu')
如果需要同时保存优化器状态(比如后续要恢复训练),也可以用同样的方式处理:
torch.save({ 'model_state_dict': model.state_dict(), 'optimizer_state_dict': optimizer.state_dict(), }, 'full_checkpoint.pth', map_location='cpu')
内容的提问来源于stack exchange,提问作者王泽君
相关产品推荐
相关产品推荐

