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

从零训练SRGAN方式是否正确?生成器出噪声损失卡滞求助

SRGAN训练故障排查:生成器输出噪声、判别器损失停滞修复

你的代码存在多处逻辑、数值、流程错误,直接导致训练崩溃,具体问题点如下:

  • 核心数值错误:Python中^是按位异或运算符,不是幂运算,你写的(10^-3)计算结果为-9,并非预期的0.001,直接导致对抗损失权重完全错乱,损失量级失衡,是训练崩溃的最主要原因。
  • 损失权重重复计算:你已经给adv_loss乘了1e-3系数,后续又乘了错误的异或结果系数,进一步放大了损失权重的错误。
  • 变量命名混乱:计算判别器损失时,真实样本损失被命名为loss_gen,生成样本损失被命名为loss_real,极易引发逻辑错误;判别器总损失未做平均,损失量级是正常值的2倍。
  • 损失实例化位置不合理:BCEWithLogitsLoss被放在批次循环内部反复初始化,产生不必要的开销。
  • 日志打印错误:用%d整数格式打印浮点型损失值,会直接截断小数部分,根本无法观测损失的细微变化,会误判损失停滞。
  • 训练流程缺失:SRGAN标准训练流程需要先用MSE损失预训练生成器,让生成器掌握基础的超分重建能力后,再开启对抗训练;直接从随机初始化状态开始对抗训练,极易出现模式崩塌,输出全噪声。
  • 缺失训练稳定性措施:没有梯度裁剪,训练初期极易出现梯度爆炸;未对齐VGG损失的输入值域,预训练VGG要求输入做ImageNet归一化,直接喂原始像素会导致内容损失梯度方向完全错误。
修正后可运行的训练代码
import torch
import torch.nn as nn

# 所有损失函数放到循环外部初始化,避免重复开销
bce_loss = nn.BCEWithLogitsLoss()
mse_loss = nn.MSELoss()

def train_model(gen, disc, vgg_loss, opt_gen, opt_disc, train_loader, device, epochs=200, pretrain_epochs=10):
    # 第一步:生成器预训练,关闭判别器梯度,只用MSE损失让生成器学会基础超分
    for p in disc.parameters():
        p.requires_grad = False
    for epoch in range(pretrain_epochs):
        for low_res, high_res in train_loader:
            low_res = low_res.to(device, non_blocking=True, dtype=torch.float).unsqueeze(1)
            high_res = high_res.to(device, non_blocking=True, dtype=torch.float).unsqueeze(1)
            gen_out = gen(low_res)
            pretrain_loss = mse_loss(gen_out, high_res)
            opt_gen.zero_grad()
            pretrain_loss.backward()
            nn.utils.clip_grad_norm_(gen.parameters(), 1.0)
            opt_gen.step()
    # 预训练完成后开启判别器梯度,进入对抗训练阶段
    for p in disc.parameters():
        p.requires_grad = True

    for epoch in range(epochs):
        run_loss_disc = 0.0
        run_loss_gen = 0.0
        for data in train_loader:
            low_res, high_res = data[0].to(device, non_blocking=True, dtype=torch.float).unsqueeze(1),\
                                data[1].to(device, non_blocking=True, dtype=torch.float).unsqueeze(1)
            
            # 训练判别器
            gen_image = gen(low_res)
            disc_fake = disc(gen_image.detach())
            disc_real = disc(high_res)
            # 修正变量命名,真实样本标签为1,生成样本标签为0,损失取平均平衡量级
            loss_disc_real = bce_loss(disc_real, torch.ones_like(disc_real))
            loss_disc_fake = bce_loss(disc_fake, torch.zeros_like(disc_fake))
            loss_disc = (loss_disc_real + loss_disc_fake) / 2

            opt_disc.zero_grad()
            loss_disc.backward()
            nn.utils.clip_grad_norm_(disc.parameters(), max_norm=1.0)
            opt_disc.step()
            run_loss_disc += loss_disc.item()

            # 训练生成器
            disc_fake = disc(gen_image)
            # 注意:传入vgg_loss前需要将gen_image、high_res对齐到预训练VGG的输入要求:值域[0,1]后做ImageNet均值方差归一化
            cont_loss = vgg_loss(gen_image, high_res)
            adv_loss = bce_loss(disc_fake, torch.ones_like(disc_fake))
            # 修正系数错误,删除错误的异或运算,统一对抗损失权重为1e-3
            gen_loss = cont_loss + 1e-3 * adv_loss

            opt_gen.zero_grad()
            gen_loss.backward()
            nn.utils.clip_grad_norm_(gen.parameters(), max_norm=1.0)
            opt_gen.step()
            run_loss_gen += gen_loss.item()

        # 修正打印格式,计算每个epoch的平均损失,保留4位小数观测变化
        avg_disc_loss = run_loss_disc / len(train_loader)
        avg_gen_loss = run_loss_gen / len(train_loader)
        print(f"Epoch [{epoch+1}/{epochs}] | Avg Disc Loss: {avg_disc_loss:.4f} | Avg Gen Loss: {avg_gen_loss:.4f}")
额外排查项
  • 检查判别器最后一层:如果用BCEWithLogitsLoss,判别器最后一层不要加Sigmoid激活,否则会导致梯度饱和,判别器无法训练。
  • 检查数据pipeline:确认低分辨率、高分辨率图像配对正确,归一化逻辑统一,不要出现值域异常(比如像素值超出[-1,1]或[0,255]范围)。
  • 检查优化器参数:SRGAN一般用Adam优化器,学习率设为1e-4~2e-4,betas设为(0.9, 0.999),学习率过大会直接导致训练崩溃。
  • 平衡生成器、判别器训练进度:如果判别器很快达到接近0的损失(分类准确率100%),生成器将无法获得有效梯度,此时可以调整训练频率,比如每训练2次生成器,训练1次判别器,避免判别器过强。
  • 检查网络结构:生成器使用残差块+PixelShuffle上采样,不要用转置卷积产生棋盘伪影;判别器使用步长卷积下采样,不要用池化层丢失高频细节。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 02:46:01