如何利用Accelerate解决PyTorch连续训练阶段GPU内存未释放问题?
解决方案
核心问题分析
使用accelerate.prepare()后,模型、优化器、数据加载器会被封装为分布式/混合精度兼容的对象,单纯del和empty_cache无法清理这些对象关联的所有GPU张量(比如优化器状态、梯度缓存、分布式通信残留张量);另外你复用的fp_model仍持有第一阶段训练的全部中间状态,是内存占用的主要来源。
具体修复步骤
清理第一阶段所有训练关联对象
确保train_loop返回训练过程中创建的优化器、学习率调度器等对象,然后一并删除:# 修改train_loop,返回优化器、调度器 fp_model, fp_loss, fp_optimizer, fp_scheduler = train_loop(model, fp_data_loader) fp_model.module.save_pretrained(checkpoint_dir) # 彻底删除所有第一阶段对象 del model, fp_data_loader, fp_optimizer, fp_scheduler先移回CPU再删除fp_model
避免模型张量因GPU引用未释放导致内存残留:# 将模型移到CPU fp_model = fp_model.cpu() del fp_model重新加载模型而非复用训练后的对象
训练后的模型会携带大量训练缓存(如梯度、动量),重新加载能完全重置状态:# 重新加载模型 gc.collect() torch.cuda.empty_cache() fp_model = AutoModelForCausalLM.from_pretrained(checkpoint_dir) # 再用accelerate重新封装 from accelerate import Accelerator accelerator = Accelerator() fp_model, sft_data_loader = accelerator.prepare(fp_model, sft_data_loader)延迟创建第二阶段数据加载器
避免提前加载数据占用GPU内存:# 移除初始的sft_data_loader = data_loader(sft_dataset) # 第一阶段清理完成后再创建 gc.collect() torch.cuda.empty_cache() sft_data_loader = data_loader(sft_dataset)分布式环境下确保所有进程同步清理
在清理后加入进程同步,避免部分进程内存未释放:from accelerate import Accelerator accelerator = Accelerator() accelerator.wait_for_everyone() gc.collect() torch.cuda.empty_cache()
完整修改后代码示例
from transformers import AutoModelForCausalLM from accelerate import Accelerator accelerator = Accelerator() # 仅初始化第一阶段数据加载器 model = AutoModelForCausalLM.from_pretrained(args.model) fp_data_loader = data_loader(fp_dataset) # 第一阶段训练:封装模型、数据加载器 model, fp_data_loader = accelerator.prepare(model, fp_data_loader) fp_model, fp_loss, fp_optimizer, fp_scheduler = train_loop(model, fp_data_loader) # 保存模型 accelerator.save(fp_model.module.state_dict(), checkpoint_dir + "/pytorch_model.bin") fp_model.module.config.save_pretrained(checkpoint_dir) # 第一阶段清理:同步进程→移模型到CPU→删除所有关联对象 accelerator.wait_for_everyone() fp_model = fp_model.cpu() del model, fp_data_loader, fp_optimizer, fp_scheduler, fp_model gc.collect() torch.cuda.empty_cache() accelerator.wait_for_everyone() # 再次同步确保所有进程完成清理 # 第二阶段初始化:重新加载模型+创建数据加载器 fp_model = AutoModelForCausalLM.from_pretrained(checkpoint_dir) sft_data_loader = data_loader(sft_dataset) # 重新封装后启动第二阶段训练 fp_model, sft_data_loader = accelerator.prepare(fp_model, sft_data_loader) sft_model, sft_loss = train_loop(fp_model, sft_data_loader)
内容的提问来源于stack exchange,提问作者hjsg1010
相关产品推荐
相关产品推荐

