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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 17:22:18