PyTorch自编码器需retain_graph=True致OOM,求优化方案
解决PyTorch自编码器训练中无需
retain_graph=True的问题 你当前遇到的问题是:不添加retain_graph=True会触发反向传播报错,但添加后又会导致显存不足。核心原因是计算图没有被正确迭代清理,以下是具体修改方案和代码示例:
关键修改方向
- 规范模型激活层使用:将函数式的
relu()改为模型成员的nn.ReLU()层,避免计算图跟踪异常 - 修正训练循环逻辑:确保每次迭代的计算图独立,移除
retain_graph=True并完善分布式训练必要步骤 - 修正损失函数注释错误:修正原注释中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()
修改说明
- 模型激活层调整:改用
nn.ReLU()层而非函数式relu(),让PyTorch更清晰地跟踪计算图的生命周期,避免函数式API可能导致的计算图残留 - 分布式训练完善:添加Horovod初始化、参数广播和优化器包装,这是分布式训练的必要步骤,缺失可能导致计算图异常
- 计算图自动清理:每次迭代的
batch都是新张量,backward()后会自动清理当前计算图的中间张量,无需手动保留
内容的提问来源于stack exchange,提问作者Eddy
相关产品推荐
相关产品推荐

