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

PyTorch实现GAN遇原地操作引发RuntimeError问题求助

问题分析与解决方案

核心问题根源

报错提示的梯度变量被原地修改,主要由以下几个问题共同导致:

  1. 损失函数的错误使用:F.binary_cross_entropy_with_logits要求输入是未经过sigmoid的原始logits,但代码中对判别器输出做了额外的sigmoid操作,既导致数值不稳定,也干扰了梯度计算流程。
  2. 判别器输出形状异常:Discriminator最后一层的Reshape()未传入任何参数,导致输出张量形状被错误修改,引发视图操作的潜在冲突。
  3. 模型与输入设备不匹配:初始化的Generator和Discriminator未移动到指定设备(CPU/GPU),混合设备计算会引发梯度异常。
  4. 训练流程的计算图冲突:在同一个批次中,先更新判别器参数,再复用之前计算的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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 18:13:09