PyTorch实现GAN遇原地操作引发RuntimeError问题求助
问题分析与解决方案
核心问题根源
报错提示的梯度变量被原地修改,主要由以下几个问题共同导致:
- 损失函数的错误使用:
F.binary_cross_entropy_with_logits要求输入是未经过sigmoid的原始logits,但代码中对判别器输出做了额外的sigmoid操作,既导致数值不稳定,也干扰了梯度计算流程。 - 判别器输出形状异常:Discriminator最后一层的
Reshape()未传入任何参数,导致输出张量形状被错误修改,引发视图操作的潜在冲突。 - 模型与输入设备不匹配:初始化的Generator和Discriminator未移动到指定设备(CPU/GPU),混合设备计算会引发梯度异常。
- 训练流程的计算图冲突:在同一个批次中,先更新判别器参数,再复用之前计算的
g_loss反向传播,此时判别器参数已被原地修改,导致梯度计算依赖的旧参数版本丢失。
分步解决方案
1. 修正损失函数
移除多余的sigmoid操作,直接使用判别器输出的logits计算损失:
def loss_nonsaturating(d, g, x_real, *, device): z = torch.randn(x_real.shape[0], g.z_dim, device=device) gz = g(z) # 直接使用判别器输出的logits,不做sigmoid dgz_logits = d(gz) dx_logits = d(x_real) # 生成与logits形状匹配的标签 real_label = torch.ones_like(dx_logits) fake_label = torch.zeros_like(dgz_logits) bce_loss = F.binary_cross_entropy_with_logits g_loss = bce_loss(dgz_logits, real_label).mean() d_loss = bce_loss(dx_logits, real_label).mean() + bce_loss(dgz_logits, fake_label).mean() return d_loss, g_loss
2. 修复Discriminator的输出层
去掉最后一层无意义的Reshape(),保留Linear层的原始输出形状:
class Discriminator(torch.nn.Module): def __init__(self, num_channels=1): super().__init__() self.net = nn.Sequential( nn.Conv2d(in_channels=1, out_channels=32, kernel_size=4, padding=1, stride=2), nn.ReLU(), nn.Conv2d(in_channels=32, out_channels=64, kernel_size=4, padding=1, stride=2), nn.ReLU(), Reshape(64*7*7), nn.Linear(64*7*7, 512), nn.ReLU(), nn.Linear(512, 1) # 移除最后的Reshape() ) def forward(self, x): return self.net(x)
3. 确保模型与输入设备一致
初始化模型时直接移动到指定设备:
g = Generator(z_dim=64).to(device) d = Discriminator().to(device)
4. 调整训练流程,避免计算图冲突
分开计算判别器和生成器的损失,训练生成器时重新计算假样本的判别器输出,避免复用已失效的计算图:
with tqdm(total=int(iter_max)) as pbar: for idx, (x, y) in enumerate(train_loader): if idx >= iter_max: break x_real, y_real = build_input(x, y, device) # 训练判别器 d_optimizer.zero_grad() z = torch.randn(x_real.shape[0], g.z_dim, device=device) gz = g(z) dx_logits = d(x_real) dgz_logits = d(gz) real_label = torch.ones_like(dx_logits) fake_label = torch.zeros_like(dgz_logits) d_loss = F.binary_cross_entropy_with_logits(dx_logits, real_label).mean() + F.binary_cross_entropy_with_logits(dgz_logits, fake_label).mean() d_loss.backward() d_optimizer.step() # 训练生成器 g_optimizer.zero_grad() # 重新生成噪声和假样本,避免复用判别器更新前的计算图 z = torch.randn(x_real.shape[0], g.z_dim, device=device) gz = g(z) dgz_logits = d(gz) g_loss = F.binary_cross_entropy_with_logits(dgz_logits, real_label).mean() g_loss.backward() g_optimizer.step() pbar.update(1)
内容的提问来源于stack exchange,提问作者Zahra Reyhanian
相关产品推荐
相关产品推荐

