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

多次调用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 17:45:29