从零训练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
相关产品推荐
相关产品推荐

