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

4GB显存GPU运行VGG16迁移学习出现CUDA显存不足如何解决?

4GB显存运行VGG16迁移学习的解决方案

通用优化思路

  • 降低batch size:当前你使用的batch size为64,是显存占用过高的核心原因,直接下调到8~16即可适配4GB显存,是改造成本最低的方案
  • 开启混合精度训练:使用PyTorch原生AMP工具将模型运算从FP32转为混合FP16/FP32,可降低近50%显存占用,同时提升训练速度
  • 梯度累加:如果小batch size影响训练效果,可以通过梯度累加的方式,在不增加显存占用的前提下等效实现大batch size的训练效果
  • 梯度检查点:如果仍有显存缺口,可以开启PyTorch梯度检查点功能,通过牺牲部分计算速度换取显存占用降低

代码修改方案

你需要对现有代码做以下修改:

  1. 修改DataLoader初始化参数,将batch_size从64下调为8
  2. 新增混合精度训练逻辑
  3. 可选新增梯度累加逻辑,等效恢复batch size为32/64的训练效果

修改后的完整代码如下:

模型定义部分

import torch
import torch.nn as nn
import torchvision
from torch.cuda.amp import GradScaler, autocast

model = torchvision.models.vgg16(pretrained=True)          
for p in model.parameters():
    p.requires_grad = False
# 新增:冻住的backbone切换为eval模式,减少显存占用和不必要的BN参数更新
model.features.eval()

sin = model.classifier[0].in_features    
model.classifier = nn.Sequential(
    nn.Linear(sin, 128), nn.ReLU(),
    nn.Linear(128, 2)
)

训练逻辑部分

DEVICE = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") 
model = model.to(DEVICE)
adam = optim.AdamW(model.parameters(), lr=1e-4, betas=(0.9, 0.99), weight_decay=2e-4)
loss_fn = nn.CrossEntropyLoss()
# 新增:混合精度缩放器
scaler = GradScaler()
# 梯度累加步数,8 * 8 = 等效batch size 64,可根据需求调整
accumulation_steps = 8

model.train()
# 注意:train模式仅作用于未冻结的分类头
for ii, (x, y) in enumerate(trainDL):
    x = x.to(DEVICE)
    y = y.to(DEVICE)
    
    # 新增:混合精度前向
    with autocast():
        z = model(x)
        loss = loss_fn(z, y)
    # 梯度归一化
    loss = loss / accumulation_steps

    # 反向传播
    scaler.scale(loss).backward()

    # 每accumulation_steps步更新一次参数
    if (ii + 1) % accumulation_steps == 0:
        scaler.step(adam)
        scaler.update()
        adam.zero_grad()

可选优化

如果调整batch size到8后仍有显存不足,可以再降低输入图像分辨率,比如将预处理的Resize参数从224调整为192,可进一步降低显存占用。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.25 20:45:11