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

PyTorch自编码器需retain_graph=True致OOM,求优化方案

解决PyTorch自编码器训练中无需retain_graph=True的问题

你当前遇到的问题是:不添加retain_graph=True会触发反向传播报错,但添加后又会导致显存不足。核心原因是计算图没有被正确迭代清理,以下是具体修改方案和代码示例:


关键修改方向

  1. 规范模型激活层使用:将函数式的relu()改为模型成员的nn.ReLU()层,避免计算图跟踪异常
  2. 修正训练循环逻辑:确保每次迭代的计算图独立,移除retain_graph=True并完善分布式训练必要步骤
  3. 修正损失函数注释错误:修正原注释中Dice公式的笔误

修改后的模型类

import torch
import torch.nn as nn
from torch.utils.data import DataLoader, DistributedSampler
import torch.optim as optim
import horovod.torch as hvd

class Autoencoder(nn.Module):
    def __init__(self, input_shape, model_config):
        super().__init__()

        output_features = model_config["output_features"]
        encode2_size = model_config["encode2_size"]
        encode3_size = model_config["encode3_size"]

        # 编码层+激活层(改为模型成员,计算图跟踪更清晰)
        self.encode1 = nn.Linear(input_shape, output_features)
        self.relu1 = nn.ReLU()
        self.encode2 = nn.Linear(output_features, encode2_size)
        self.relu2 = nn.ReLU()
        self.encode3 = nn.Linear(encode2_size, encode3_size)
        self.relu3 = nn.ReLU()

        # 解码层+激活层
        self.decode1 = nn.Linear(encode3_size, encode2_size)
        self.relu4 = nn.ReLU()
        self.decode2 = nn.Linear(encode2_size, output_features)
        self.relu5 = nn.ReLU()
        self.decode3 = nn.Linear(output_features, input_shape)
        self.relu6 = nn.ReLU()

    def encode(self, x: torch.Tensor):
        x = self.relu1(self.encode1(x))
        x = self.relu2(self.encode2(x))
        x = self.relu3(self.encode3(x))
        return x

    def decode(self, x: torch.Tensor):
        x = self.relu4(self.decode1(x))
        x = self.relu5(self.decode2(x))
        x = self.relu6(self.decode3(x))
        return x

    def forward(self, x: torch.Tensor):
        x = self.encode(x)
        x = self.decode(x)
        return x

修改后的损失函数

class DiceLoss(nn.Module):
    """
    The formula: 2*|X ∩ Y|/(|X| + |Y|)
    """

    def __init__(self, weight=None, size_average=True):
        super(DiceLoss, self).__init__()

    def forward(self, inputs, targets, smooth=1):
        # comment out if your model contains a sigmoid activation
        inputs = torch.sigmoid(inputs)

        # flatten label and prediction tensors
        inputs = inputs.view(-1)
        targets = targets.view(-1)

        intersection = (inputs * targets).sum()
        dice = (2. * intersection + smooth) / (inputs.sum() + targets.sum() + smooth)

        return 1 - dice

修改后的训练代码

# 初始化Horovod(分布式训练必要步骤)
hvd.init()

train_sampler = DistributedSampler(data_tensor, num_replicas=hvd.size(), rank=hvd.rank())
train_dataloader = DataLoader(data_tensor, batch_size=batch_size, shuffle=False, sampler=train_sampler)

# autoencoder params
epochs = model_config["epochs"]
net = Autoencoder(embedding_dim, model_config)
# 广播模型参数到所有进程
net = hvd.broadcast_parameters(net, root_rank=0)

# loss function and optimizer
loss_function = DiceLoss()
optimizer = optim.Adagrad(net.parameters(), lr=model_config["lr"], weight_decay=model_config["weight_decay"])
# 用Horovod包装优化器
optimizer = hvd.DistributedOptimizer(optimizer, named_parameters=net.named_parameters())

for epoch in range(epochs):
    train_sampler.set_epoch(epoch)  # 分布式训练确保每个epoch采样不重复
    net.train()
    for i, batch in enumerate(train_dataloader):
        net.zero_grad()
        # 将batch移到对应设备(GPU/CPU)
        batch = batch.to(hvd.local_rank())
        # 前向传播
        output = net(batch)
        # 计算损失
        loss = loss_function(output, batch)
        # 反向传播(移除retain_graph=True)
        loss.backward()
        optimizer.step()

修改说明

  1. 模型激活层调整:改用nn.ReLU()层而非函数式relu(),让PyTorch更清晰地跟踪计算图的生命周期,避免函数式API可能导致的计算图残留
  2. 分布式训练完善:添加Horovod初始化、参数广播和优化器包装,这是分布式训练的必要步骤,缺失可能导致计算图异常
  3. 计算图自动清理:每次迭代的batch都是新张量,backward()后会自动清理当前计算图的中间张量,无需手动保留

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.26 05:17:35