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

PyTorch中使用自定义损失函数时如何实现内存高效反向传播?

PyTorch大规模项目反向传播内存优化(自定义损失场景)

针对你在大规模深度学习项目中遇到的反向传播内存占用过高问题,结合你提供的自定义损失函数,以下是几种不影响损失计算结果的内存高效优化方案:

1. 优化自定义损失的计算逻辑,避免冗余中间张量

你的当前损失计算torch.sum(x * y)会先生成一个与x/y同维度的乘积张量,再求和。这个中间乘积张量会占用大量内存,且反向传播时需保留其梯度信息。可以直接用等价计算方式跳过中间张量的生成:

class CustomLoss(nn.Module):
    def forward(self, x, y):
        # 等价于sum(x*y),但不生成大尺寸中间张量
        return torch.dot(x.flatten(), y.flatten())
        # 若x/y维度不规则,也可用einsum实现任意维度的点积求和
        # return torch.einsum('...,...->', x, y)

2. 梯度检查点(Gradient Checkpointing)

使用PyTorch内置的torch.utils.checkpoint,在正向传播时不保存所有中间激活张量,反向传播时按需重新计算必要的激活,以此大幅降低内存占用。适用于网络结构较深或输入尺寸极大的场景:

方案A:给自定义损失添加检查点

from torch.utils.checkpoint import checkpoint

class CustomLoss(nn.Module):
    def forward(self, x, y):
        # 用checkpoint包裹损失计算逻辑
        def loss_func(x, y):
            return torch.dot(x.flatten(), y.flatten())
        return checkpoint(loss_func, x, y)

方案B:给网络层添加检查点

如果内存压力主要来自网络前向传播的激活,可以在Net的前向传播中对部分层使用检查点:

from torch.utils.checkpoint import checkpoint

class Net(nn.Module):
    def __init__(self):
        super().__init__()
        self.layer1 = nn.Linear(1024, 2048)
        self.layer2 = nn.Linear(2048, 1024)
    
    def forward(self, x):
        # 对layer2的前向传播使用检查点
        x = self.layer1(x)
        x = checkpoint(self.layer2, x)
        return x

3. 混合精度训练

通过自动混合精度(AMP)将部分张量从FP32转为FP16存储,能直接减少约50%的内存占用,且几乎不影响模型精度:

from torch.cuda.amp import GradScaler, autocast

# 初始化混合精度工具
scaler = GradScaler()
model = Net().cuda()
loss_fn = CustomLoss().cuda()
optimizer = torch.optim.Adam(model.parameters())

# 训练循环中使用混合精度
for x, y in dataloader:
    x, y = x.cuda(), y.cuda()
    optimizer.zero_grad()
    
    # 正向传播用autocast自动转换精度
    with autocast():
        output = model(x)
        loss = loss_fn(output, y)
    
    # 反向传播用scaler处理梯度缩放
    scaler.scale(loss).backward()
    scaler.step(optimizer)
    scaler.update()

4. 手动清理冗余张量

在自定义损失的前向传播中,及时删除不再需要的中间张量,并清理CUDA缓存(仅GPU场景):

class CustomLoss(nn.Module):
    def forward(self, x, y):
        product = x * y
        loss = torch.sum(product)
        # 删除中间张量,减少内存引用
        del product
        # 仅在使用GPU时调用,清理未被引用的显存
        if torch.cuda.is_available():
            torch.cuda.empty_cache()
        return loss

5. 限制反向传播的梯度范围

如果你的模型存在梯度爆炸或过大的梯度张量,可以通过torch.nn.utils.clip_grad_norm_限制梯度范数,间接减少梯度张量的内存占用:

# 反向传播后添加梯度裁剪
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.08 05:43:13