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

实现NetVLAD层Triplet Loss时CUDA内存不足问题及优化咨询

NetVLAD Triplet Loss 显存不足问题优化方案

原代码核心问题

  1. 负样本加载逻辑错误:next(iter(train_loader))会重复创建数据加载器迭代器,打乱训练流程的同时额外占用资源,还可能导致数据重复读取。
  2. 显存未及时释放:锚点、正样本、负样本三个批次及模型前向传播的中间张量同时驻留显存,快速耗尽CUDA内存。

具体优化方案

1. 修复负样本加载逻辑

使用持久化的负样本迭代器获取随机批次,避免重复创建迭代器:

# 提前初始化负样本迭代器
neg_iter = iter(train_loader)

for i, (inputs, poses) in enumerate(train_loader):
    inputs = inputs.to(device)
    positives = positive_tfs(inputs)
    positives = positives.to(device)
    
    # 获取负样本,迭代器耗尽时重新初始化
    try:
        negatives, _ = next(neg_iter)
    except StopIteration:
        neg_iter = iter(train_loader)
        negatives, _ = next(neg_iter)
    negatives = negatives.to(device)
    
    # 后续训练逻辑...

2. 主动释放显存,减少同时驻留的张量

计算完无用张量后立即删除,必要时清理显存缓存:

optimizer.zero_grad()

# 计算锚点VLAD,释放无用输入和中间张量
a_vlad, pos_out, ori_out = model(inputs)
del pos_out, ori_out, inputs
torch.cuda.empty_cache()

# 计算正样本VLAD,释放正样本输入
p_vlad = model(positives, get_pose=False)
del positives
torch.cuda.empty_cache()

# 计算负样本VLAD,释放负样本输入
n_vlad = model(negatives, get_pose=False)
del negatives
torch.cuda.empty_cache()

# 计算损失后释放嵌入张量
loss = triplet_loss(a_vlad, p_vlad, n_vlad)
del a_vlad, p_vlad, n_vlad
torch.cuda.empty_cache()

# 反向传播与优化
loss.backward()
optimizer.step()
del loss
torch.cuda.empty_cache()

3. 开启混合精度训练

借助PyTorch自动混合精度减少显存占用,几乎不影响模型性能:

from torch.cuda.amp import GradScaler, autocast

# 初始化混合精度工具
scaler = GradScaler()
neg_iter = iter(train_loader)

for i, (inputs, poses) in enumerate(train_loader):
    inputs = inputs.to(device)
    positives = positive_tfs(inputs)
    positives = positives.to(device)
    
    try:
        negatives, _ = next(neg_iter)
    except StopIteration:
        neg_iter = iter(train_loader)
        negatives, _ = next(neg_iter)
    negatives = negatives.to(device)
    
    optimizer.zero_grad()
    
    # 开启自动混合精度上下文
    with autocast():
        a_vlad, pos_out, ori_out = model(inputs)
        del pos_out, ori_out, inputs
        p_vlad = model(positives, get_pose=False)
        del positives
        n_vlad = model(negatives, get_pose=False)
        del negatives
        loss = triplet_loss(a_vlad, p_vlad, n_vlad)
        del a_vlad, p_vlad, n_vlad
    
    # 混合精度反向传播
    scaler.scale(loss).backward()
    scaler.step(optimizer)
    scaler.update()
    
    del loss
    torch.cuda.empty_cache()

4. 降低批次大小

如果以上方法仍无法解决显存问题,直接调小train_loader的batch_size(比如从64改为32或16),这是最快速的显存减压方式。


内容的提问来源于stack exchange,提问作者Shania F.

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 15:52:44