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

如何利用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 17:04:52