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
相关产品推荐
相关产品推荐

