多次调用VGG提取图像特征引发OOM错误,求原因排查与解决方法
解决VGG特征提取时的OOM内存崩溃问题
我需要提取8091张图像的VGG特征,将结果以形状为(1,4096)的张量形式存在内存字典里,但刚完成约6%就因内存不足(OOM)崩溃。测试发现不是单纯内存空间不够——哪怕只存VGG的分类结果也会触发错误,而用随机生成同等尺寸张量的方式完全没问题。
复现错误的最简代码
import torch, torchvision from tqdm import tqdm vgg = torchvision.models.vgg16(weights='DEFAULT') def try_and_crash(gen_data): store_out = {} for i in tqdm(range(8091)): my_output = gen_data(torch.randn(1,3,224,224)) store_out[i] = my_output return store_out
对比测试
- 调用随机张量生成函数,运行完全正常:
just_fine = try_and_crash(lambda x: torch.randn(1,4096))
- 调用VGG则直接触发内存崩溃:
will_crash = try_and_crash(vgg)
问题原因
核心问题是PyTorch默认会保留计算图用于反向传播。每次调用VGG时,输出张量my_output会附带整个前向传播过程的计算图引用,这些计算图占用的内存远超过张量本身的存储大小,累积到一定程度就会触发OOM。而随机生成的张量没有计算图关联,内存占用仅为张量本身的大小,因此不会出现问题。
解决方案
1. 禁用计算图跟踪
使用torch.no_grad()上下文管理器包裹模型调用,彻底关闭计算图的生成和保留,这是解决问题的关键:
def try_and_fix(gen_data): store_out = {} for i in tqdm(range(8091)): with torch.no_grad(): my_output = gen_data(torch.randn(1,3,224,224)) store_out[i] = my_output return store_out # 现在调用VGG不会崩溃 fixed = try_and_fix(vgg)
2. 切换模型到评估模式
将模型设置为评估模式,关闭训练相关的特殊层(如Dropout、BatchNorm),进一步优化内存占用和计算速度:
vgg = torchvision.models.vgg16(weights='DEFAULT') vgg.eval() # 启用评估模式
3. 可选:将输出张量转到CPU存储
如果GPU内存仍然紧张,可以在计算完成后将张量转移到CPU内存存储,释放GPU空间:
with torch.no_grad(): my_output = gen_data(torch.randn(1,3,224,224)).cpu()
内容的提问来源于stack exchange,提问作者Ferdinando Randisi
相关产品推荐
相关产品推荐

