实现NetVLAD层Triplet Loss时CUDA内存不足问题及优化咨询
NetVLAD Triplet Loss 显存不足问题优化方案
原代码核心问题
- 负样本加载逻辑错误:
next(iter(train_loader))会重复创建数据加载器迭代器,打乱训练流程的同时额外占用资源,还可能导致数据重复读取。 - 显存未及时释放:锚点、正样本、负样本三个批次及模型前向传播的中间张量同时驻留显存,快速耗尽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.
相关产品推荐
相关产品推荐

