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

Colab加载Google Drive中torch.save保存的GAN内存超运行时40倍原因求解

问题原因
  • 保存逻辑问题:你当前采用的是直接保存整个GAN实例对象的方式,而非仅保存模型状态字典state_dict。PyTorch序列化整个模型对象时,会附带保存训练过程中生成的计算图、梯度缓存、未被清理的中间张量等冗余数据,这些数据在运行时会被PyTorch自动回收释放,因此运行时内存占用低,但序列化后会永久保存在文件中,加载时会全部载入内存;如果你的GAN类中绑定了数据集、数据加载器等大对象的引用,这些内容也会被一并序列化,是内存暴增的核心原因。
  • 加载逻辑冗余:你的loadPopulation函数中定义了popDictArr列表持有所有加载的字典对象,提取出GAN实例后没有释放该列表的引用,字典中附带的大量冗余数据无法被Python垃圾回收机制清理,进一步占用内存。
  • 加载参数缺失:调用torch.load时未指定map_location参数,若你保存的是GPU训练后的模型版本,新会话加载时会先将所有模型权重加载到CPU内存,若后续转移到GPU还会额外生成一份副本,额外占用内存。
修复方案
  • 优先调整保存逻辑,仅保存模型的state_dict而非整个实例,从根源避免冗余数据被序列化:
# 修改后的保存代码示例
def saveEntirePopulation(keyPath, population):
    for ind, gan in enumerate(population):
        torch.save({
            'generator_state': gan.generator.state_dict(),
            'discriminator_state': gan.discriminator.state_dict(),
            # 如有优化器、超参数等需要保留的内容可在此处追加
        }, keyPath + 'population_' + str(ind))
  • 调整加载逻辑,加载后及时释放冗余引用,指定map_location避免多余副本:
def loadPopulation(keyPath, popSize, device='cuda' if torch.cuda.is_available() else 'cpu'):
    popArr = []
    for ind in range(popSize):
        # 直接加载到目标设备,避免中间副本
        ckpt = torch.load(keyPath + 'population_' + str(ind), map_location=device)
        # 替换为你自己的GAN类初始化代码
        gan = GAN()
        gan.generator.load_state_dict(ckpt['generator_state'])
        gan.discriminator.load_state_dict(ckpt['discriminator_state'])
        gan.generator.to(device)
        gan.discriminator.to(device)
        gan.generator.train()
        gan.discriminator.train()
        popArr.append(gan)
        # 手动删除检查点对象,清空引用
        del ckpt
    # 手动触发垃圾回收清理残留
    import gc
    gc.collect()
    return popArr
  • 如果你确实需要保存整个模型对象,保存前先调用gan.generator.zero_grad()、gan.discriminator.zero_grad()清理梯度缓存,同时删除模型中绑定的数据集、中间张量等不需要的属性后再进行序列化。

内容的提问来源于stack exchange,提问作者GammaRamma

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.05 03:24:04