pytorch中gan的generator和discriminator需严格对齐数据维度、通道数、尺寸与归一化策略:generator输出通道ch必须等于数据集通道数(如mnist为1、cifar-10为3),最后一层用tanh或sigmoid;discriminator输入尺寸和通道数须与generator输出完全一致,并在forward开头加assert校验;batchnorm2d在generator中禁用于输出层,在discriminator中避开第一层;权重须用normal_(0, 0.02)初始化,且训练中batchnorm2d必须保持training=true。

PyTorch里GAN的Generator和Discriminator不是靠“模板”拼出来的,而是得按数据维度、激活方式、归一化策略一层层对齐输入输出——错一个shape或一处nn.BatchNorm2d的位置,训练就直接崩成NaN。
Generator输出通道数必须严格匹配真实图像的channel数
比如你训的是CIFAR-10(3通道RGB),Generator最后一层的out_channels就得是3;如果是MNIST(单通道灰度图),就必须是1。常见错误是直接抄代码把out_channels=3写死,结果喂MNIST时生成全黑图——因为输出被sigmoid压到[0,1]后,3通道叠加导致像素值溢出或均值偏移。
- 检查你的数据加载器:
train_loader.dataset[0][0].shape看清楚是(1, 28, 28)还是(3, 32, 32) -
Generator最后一层用nn.Conv2d(in_channels=64, out_channels=CH, kernel_size=4, stride=2, padding=1),其中CH必须等于数据集的通道数 - 别忘了接
nn.Tanh()(对[-1,1]范围更稳)或nn.Sigmoid()(只适用于[0,1]输入且已做归一化)
Discriminator输入尺寸要和Generator输出完全一致
很多人把Discriminator写成固定接受(3, 64, 64),但Generator输出是(3, 32, 32),结果forward直接报size mismatch。这不是维度推错,是根本没对齐。
Python 3.14.2是Python编程语言在2025年12月5日发布的稳定版本,属于3.14系列的第二个维护更新。该版本包含了18项修复,重点解决了多进程、数据类及正则表达式等模块的回归问题,并修复了CVE-2025-12084等安全漏洞。此版本标志着自由线程模式(移除GIL)正式获得官方支持,是Python发展的重要里程碑。
- 先确定目标图像尺寸:CIFAR-10是
32×32,LSUN bedroom常用64×64,自定义数据集务必用torchvision.transforms.Resize统一预处理 -
Discriminator第一层nn.Conv2d(in_channels=CH, ...)的in_channels必须等于Generator输出通道,且输入tensor的H和W必须和Generator最后一层输出尺寸一致 - 推荐在
Discriminator.forward开头加一句assert x.shape[2:] == (H, W), f"expected {(H, W)}, got {x.shape[2:]}",早报错早调试
BatchNorm2d在Generator中不能放在最末层,但在Discriminator中通常要避开输入层
nn.BatchNorm2d在生成器里加在ConvTranspose2d之后能稳定训练,但若放在最后一层(即输出层前),会导致生成图像整体偏亮/偏暗;在判别器里,如果第一层就加BN,会破坏真实样本的原始分布统计量,削弱判别能力。
- Generator结构建议:
ConvTranspose2d → BatchNorm2d → ReLU(中间层),最后一层只留ConvTranspose2d → Tanh - Discriminator结构建议:第一层用
Conv2d → LeakyReLU(不加BN),中间层用Conv2d → BatchNorm2d → LeakyReLU,输出层只留Conv2d - 注意:GAN训练中
BatchNorm2d的training=True必须始终开启(哪怕在eval模式下推理生成图也要手动设.train()),否则统计量冻结会导致生成质量骤降
初始化权重比调学习率还关键
GAN极容易因权重初始太大会让梯度爆炸,或者太小导致梯度消失——尤其当用了ConvTranspose2d这种本身就不稳定的上采样层时。
- Generator所有
ConvTranspose2d和Conv2d层,用torch.nn.init.normal_(m.weight.data, 0.0, 0.02) - Discriminator所有
Conv2d层,同样用normal_(..., 0.0, 0.02);BatchNorm2d的weight设为1、bias设为0 - 千万别用默认初始化,也别用
Kaiming——GAN对初始化敏感度远超分类任务
真正卡住的往往不是网络结构,而是torch.Size没对齐、BatchNorm2d误开eval()、或者权重初始化漏了某一层。跑通第一个batch的loss不为nan,比写出“完美架构”重要十倍。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!










