PyTorch多输入图像单张梯度计算及CUDA显存不足问题求助
问题原因
训练阶段用多GPU数据并行时,每个GPU仅负责处理部分模型参数或数据分片,中间激活和梯度也被分摊到多个GPU;而当前评估阶段单GPU运行时,计算输入图像的梯度需要保留所有中间层的激活张量(反向传播依赖这些激活计算输入的梯度),加上输入张量本身的显存占用,远超单GPU的承载能力——哪怕模型参数的显存占用比输入梯度+中间激活小很多,因为中间激活的显存通常是参数的数倍。
你的输入张量[1,6,3,928,1600]单精度下本身约100MB,但模型的注意力层(报错点在image_cross_attention.py的torch.cat(attns, dim=1))会生成大量中间张量,这些张量在计算输入梯度时必须全部保留,这才是显存耗尽的核心原因。
可行解决方案
1. 启用混合精度(最有效,几乎不影响精度)
用PyTorch的自动混合精度(AMP),将中间激活用半精度(FP16)存储,大幅降低显存占用,同时不影响梯度计算的精度。修改代码如下:
from torch.cuda.amp import autocast, GradScaler my_model.eval() for param in my_model.parameters(): param.requires_grad = False torch.cuda.empty_cache() scaler = GradScaler() # 用于半精度梯度缩放 for i_iter_val, (imgs, img_metas, val_vox_label, val_grid, val_pt_labs) in enumerate(val_dataset_loader): imgs = imgs.cuda() imgs.requires_grad = True with autocast(): # 启用半精度推理 predict_labels_vox, predict_labels_pts = my_model(img=imgs, ...) loss = ... # 你的损失计算逻辑 my_model.zero_grad() scaler.scale(loss).backward() # 半精度反向传播 data_grad = imgs.grad.data # 后续扰动添加逻辑...
2. 分批计算输入图像的梯度(针对6张图像拆分)
每次仅让1张图像的梯度可求,其余图像固定,逐张计算梯度后合并,这样每次的中间激活仅对应单张图像,显存占用降低6倍左右:
my_model.eval() for param in my_model.parameters(): param.requires_grad = False torch.cuda.empty_cache() for i_iter_val, (imgs, img_metas, val_vox_label, val_grid, val_pt_labs) in enumerate(val_dataset_loader): imgs = imgs.cuda() data_grad = torch.zeros_like(imgs) # 初始化梯度存储 # 逐张处理6张图像 for idx in range(6): imgs_single = imgs.clone() imgs_single.requires_grad = False imgs_single[:, idx:idx+1, :, :, :].requires_grad = True # 仅当前图像可求梯度 predict_labels_vox, predict_labels_pts = my_model(img=imgs_single, ...) loss = ... # 你的损失计算逻辑 my_model.zero_grad() loss.backward() # 提取当前图像的梯度并保存 data_grad[:, idx:idx+1, :, :, :] = imgs_single[:, idx:idx+1, :, :, :].grad.data.clone() torch.cuda.empty_cache() # 清理当前批次的中间显存 # 后续用data_grad添加扰动...
3. 多GPU并行计算梯度
利用服务器的4块GPU,将输入拆分到多个GPU上并行计算,分摊显存压力。用DataParallel包装模型:
# 初始化多GPU from torch.nn.parallel import DataParallel my_model = DataParallel(my_model).cuda() my_model.eval() for param in my_model.parameters(): param.requires_grad = False torch.cuda.empty_cache() for i_iter_val, (imgs, img_metas, val_vox_label, val_grid, val_pt_labs) in enumerate(val_dataset_loader): imgs = imgs.cuda() imgs.requires_grad = True predict_labels_vox, predict_labels_pts = my_model(img=imgs, ...) loss = ... my_model.zero_grad() loss.backward() data_grad = imgs.grad.data # 后续扰动逻辑...
注:DataParallel会自动将输入拆分到多个GPU,每个GPU处理部分图像,中间激活和梯度都分摊到各GPU,避免单GPU显存溢出。
4. 优化显存碎片化
设置环境变量减少显存碎片化,让PyTorch更高效利用剩余显存:
export PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128
或者在Python代码开头添加:
import os os.environ['PYTORCH_CUDA_ALLOC_CONF'] = 'max_split_size_mb:128'
5. 梯度检查点(牺牲速度换显存)
对模型中显存占用大的层(比如你的image_cross_attention模块)使用梯度检查点,反向传播时重新计算这些层的激活,而不是存储:
from torch.utils.checkpoint import checkpoint # 修改image_cross_attention模块的get_sampling_offsets_and_attention方法 def get_sampling_offsets_and_attention(self, query): # 将原逻辑包装在checkpoint中 def _forward(query): # 原有的attns计算逻辑 attns = [...] attns = torch.cat(attns, dim=1) return sampling_offsets, attns return checkpoint(_forward, query)
注:梯度检查点会增加计算时间,适合显存紧张但时间充足的场景。
内容的提问来源于stack exchange,提问作者Hengwei Chen

