ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

PyTorch DCGAN二次元头像生成:数据清洗、训练与调优实战

PyTorch DCGAN二次元头像生成:数据清洗、训练与调优实战 二次元人物头像生成这个方向几乎是每个刚接触生成模型的人都会拿来练手的项目。原因很实在数据好找、效果肉眼可见、训练完之后能真的生成一批能用的图正反馈来得快。而 PyTorch 配合 GAN 这套组合又是把理论和能跑起来的代码连接得最短的一条路。这篇内容我会把整个流程从头到尾拆开讲一遍——数据集怎么清洗、网络结构为什么这么设计、GAN 的数学原理怎么一行行对应到代码、训练过程中哪些参数不能乱动、以及那些只有真正跑过几十轮实验才会知道的小坑。全文基于 DCGAN 架构输出 64×64 的二次元头像单卡消费级显卡就能跑完代码可以直接抄走改成自己的数据集。适合已经会一点 Python、装过 PyTorch、但还没完整跑通过一个生成模型的人如果你已经做过图像分类这篇也能帮你把判别式模型的思维切换到生成式模型上。1. 项目整体设计与落地思路先把要做的事情说清楚我们要训练一个网络输入一串随机噪声输出一张 64×64 的二次元风格头像。这件事的难点不在写网络而在于让生成器和判别器在对抗中保持平衡——任何一方过强训练都会崩。所以在动手写代码之前先把架构选型、参数分配和算力预算这三件事想明白后面会省下大量返工时间。1.1 为什么用 DCGAN 而不是把全连接网络硬怼上去最早的 GAN 论文里生成器和判别器都是全连接网络输入噪声直接映射到像素向量。这个做法在 MNIST 这种 28×28 的灰度图上勉强能用但一旦换成 64×64 的彩色图参数量会爆炸。算一笔账64×64×3 等于 12288 维输出如果中间藏一层 1024 维的全连接层光这一层就是 12288×1024 约 1258 万个参数而且生成出来的图基本是一团彩色的糊状物完全没有局部结构。DCGAN 的核心贡献就是把卷积结构系统地引入 GAN并且定下了几条至今仍被沿用的经验规则判别器里用步长卷积代替池化层、生成器里用转置卷积做上采样、两个网络都加批归一化生成器输出层和判别器输入层除外、生成器用 ReLU 而判别器用 LeakyReLU。这几条规则背后的逻辑是一致的——保持梯度的流动性和空间信息的连续性。池化层会丢掉位置信息而生成任务恰恰需要模型学到眼睛在上面、嘴巴在下面这种空间先验所以用带步长的卷积自己学下采样方式更合理。实际项目里还有一个现实考量DCGAN 的结构足够简单你能在一张 8GB 显存的卡上把 batch size 开到 64 甚至 128训练轮次跑得快方便反复试错。像 StyleGAN 那种动辄上千万参数的架构光是理解网络结构就要花掉一周不适合拿来入门。1.2 生成器与判别器的对称性与参数量分配一个很容易被忽略的细节是生成器和判别器的容量应该大致对等但不必完全相等。我自己的习惯是让两者的参数量处于同一量级差距控制在一倍以内。如果判别器明显更强它会很快学会分辨真假生成器的梯度就会消失反过来判别器太弱生成器就会开始骗自己输出退化成少数几种图。下面这张表是按 64×64 输出、噪声维度 128、基础通道数 64 算出来的参数量分布你可以对照着自己的配置核对一遍模块层类型输入通道输出通道卷积核参数量生成器 G1ConvTranspose2d1285124×41,048,576生成器 G2ConvTranspose2d5122564×42,097,152生成器 G3ConvTranspose2d2561284×4524,288生成器 G4ConvTranspose2d128644×4131,072生成器 G5ConvTranspose2d6434×43,072判别器 D1Conv2d3644×43,072判别器 D2Conv2d641284×4131,072判别器 D3Conv2d1282564×4524,288判别器 D4Conv2d2565124×42,097,152判别器 D5Conv2d51214×48,192生成器总计约 380 万参数判别器约 276 万参数比值 1.38属于健康范围。注意生成器的 G2 层参数量最大因为它要同时承担通道数减半和特征图翻倍两件事。如果你把基础通道数从 64 提到 128参数量会涨到接近四倍这时候 8GB 显存可能就吃紧了需要把 batch size 降到 32。1.3 一次训练要花多少资源先算账再动手在开始之前先估算一下训练成本能避免跑到一半发现显存不够或者时间不可接受。假设数据集是 5 万张 64×64 头像batch size 128那么一个 epoch 是 391 个 iteration。每个 iteration 前向加反向大致是 3 次网络传播判别器真样本、判别器假样本、生成器按单张 2080Ti 或同级显卡估算大约 0.25 秒一个 iteration一个 epoch 大概 100 秒。实际经验是DCGAN 在二次元头像上要到能看的效果通常需要 80 到 150 个 epoch也就是 2.5 到 4 小时。如果你的显卡更弱这个时间会线性拉长。所以在动手之前建议先用 5000 张图跑 10 个 epoch 做一次冒烟测试确认损失曲线形态正常、生成结果有从噪声向色块演化的趋势再上全量数据。这个习惯帮我省过很多次跑了六小时发现数据路径写错了的尴尬。提示训练日志里一定要记录每轮的生成器损失、判别器损失以及每 5 个 epoch 存一批生成样本图。事后复盘时这批样本图的价值远大于损失数值本身。2. 数据集准备从原始图片到能进网络的张量数据这一块我打算讲得细一点因为在我自己的经历里GAN 效果差的原因有七成出在数据上而不是网络结构或超参数。很多人抱怨生成的都是鬼脸最后一查要么是图片长宽比被硬拉伸了要么是数据集里混进了大量半身像甚至风景图模型被迫去拟合这些完全不同的分布。2.1 数据集从哪来、目录怎么组织二次元头像数据集的获取渠道比较固定常见做法是抓取角色立绘后做裁剪或者使用公开的头像集合。无论来源如何落盘时统一成下面这种扁平结构最省事data/ faces/ 000001.jpg 000002.jpg ... 050000.jpg只用一层目录文件名不携带语义信息。为什么不用按角色或画师分文件夹因为无条件 GAN 不关心标签多层目录只会让 DataLoader 的路径遍历逻辑变复杂。真正需要按类别训练的时候也就是条件 GAN再改成子目录结构也不迟。数量上我的建议是至少 2 万张5 万张是比较舒服的量级。少于 1 万张时判别器会极快地记住所有训练样本导致它对新样本一律判假生成器拿不到有效梯度训练很快陷入僵局。这个现象在小数据集上几乎是必然的和你的网络设计无关。2.2 清洗环节人脸检测、方形裁剪与去重清洗流程我一般分三步走。第一步是格式统一把所有图片转成 RGB 三通道这一步能干掉 RGBA 透明通道和灰度图带来的维度不一致问题否则你在 ToTensor 之后会碰到 1 通道和 3 通道混着进网络的情况报错信息还挺绕。第二步是裁剪出正方形并定位到人脸区域。这里可以用 OpenCV 自带的级联检测器做人脸框定位二次元风格虽然和人脸检测器的训练分布不完全一致但检出率通常在八成以上足够用了。检出失败的样本退化成以图像中心裁一个边长等于短边的正方形比直接丢弃要划算毕竟数据量宝贵。import cv2 from pathlib import Path from PIL import Image CASCADE cv2.CascadeClassifier( cv2.data.haarcascades haarcascade_frontalface_default.xml ) def square_crop_and_save(src_path: Path, dst_dir: Path, size: int 64): img cv2.imread(str(src_path)) if img is None: return False h, w img.shape[:2] gray cv2.cvtColor(img, cv2.COLOR_BGR2GRAY) faces CASCADE.detectMultiScale(gray, scaleFactor1.1, minNeighbors5) if len(faces) 0: # 取面积最大的一个框 x, y, fw, fh max(faces, keylambda r: r[2] * r[3]) cx, cy x fw / 2, y fh / 2 side int(max(fw, fh) * 1.35) # 外扩 35%把头发和下巴包进来 else: cx, cy w / 2, h / 2 side int(min(w, h) * 0.8) half side // 2 x1, y1 int(max(0, cx - half)), int(max(0, cy - half)) x2, y2 int(min(w, x1 side)), int(min(h, y1 side)) if x2 - x1 32 or y2 - y1 32: return False crop img[y1:y2, x1:x2] crop cv2.resize(crop, (size, size), interpolationcv2.INTER_AREA) out dst_dir / src_path.name cv2.imwrite(str(out), crop, [cv2.IMWRITE_JPEG_QUALITY, 95]) return True外扩系数取 1.35 是调出来的经验值。取 1.0 的话很多图只剩一张脸头发、耳朵这些二次元角色的辨识特征全没了生成结果会呈现出一种证件照的呆板感取 1.0 以下则会切掉额头。1.3 到 1.4 之间比较稳。第三步是去重和异常剔除。用感知哈希做粗筛汉明距离小于 5 的判为重复再用文件大小和像素方差筛掉纯色图、损坏图。这一步做完数据集里那些一看就是坏样本基本能清干净。2.3 Dataset 与 DataLoader 的落地写法自己写 Dataset 类的时候有个小细节值得注意不要在__init__里把图片全部读进内存。5 万张 64×64 的图虽然只有几百 MB但加上 OpenCV 的额外开销多进程 DataLoader 下很容易内存翻倍。正确做法是只收集路径列表在__getitem__里按需读取。from pathlib import Path from PIL import Image from torch.utils.data import Dataset IMG_EXT {.jpg, .jpeg, .png, .webp, .bmp} class AnimeFaceDataset(Dataset): def __init__(self, root: str, transformNone): self.paths sorted( p for p in Path(root).rglob(*) if p.suffix.lower() in IMG_EXT ) if len(self.paths) 0: raise RuntimeError(f没有在 {root} 下找到任何图片检查路径) self.transform transform def __len__(self): return len(self.paths) def __getitem__(self, idx): path self.paths[idx] try: img Image.open(path).convert(RGB) except Exception: # 单张图损坏时返回一张黑图保证训练不中断 img Image.new(RGB, (64, 64), (0, 0, 0)) if self.transform is not None: img self.transform(img) return img那个 try-except 的兜底看着不起眼但在真实数据集上非常必要。我遇到过一次训练跑到第 63 个 epoch 时因为一张半截损坏的 JPEG 直接崩掉没有断点续训的话等于白跑。DataLoader 的配置loader DataLoader( dataset, batch_size128, shuffleTrue, num_workers4, pin_memoryTrue, drop_lastTrue, persistent_workersTrue, )drop_lastTrue必须开。如果最后一个 batch 只有 1 张图批归一化在那一批上算出的均值和方差毫无统计意义会给判别器带来一次剧烈的参数抖动。pin_memory让数据从页内存锁到显存配合 GPU 训练能省下不少搬运时间。2.4 归一化与数据增强的取舍归一化统一用Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))也就是把像素从 [0,1] 映射到 [-1,1]。这么做是为了配合生成器最后一层的 Tanh 激活——Tanh 的值域就是 [-1,1]如果训练数据还是 [0,1]判别器光凭输出的数值范围就能轻松分辨真假训练会从一开始就失衡。这个坑我踩过当时损失曲线看着挺正常但生成图一直是灰蒙蒙的怎么调学习率都没用。数据增强在 GAN 里要谨慎。随机水平翻转是安全的因为二次元头像基本左右对称翻转不会引入虚假的分布。但随机裁剪、颜色抖动、旋转这几类操作要小心颜色抖动会让生成器学到同一张脸有多种色调输出变得浑浊小角度旋转则会引入黑边判别器很容易学会看到黑边就判假反而干扰了它对内容的判断。我的配置是只保留水平翻转加缩放from torchvision import transforms transform transforms.Compose([ transforms.Resize(72, interpolationtransforms.InterpolationMode.BICUBIC), transforms.CenterCrop(64), transforms.RandomHorizontalFlip(p0.5), transforms.ToTensor(), transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5)), ])先放大到 72 再中心裁到 64等于给了一个很小的抖动空间同时避免了直接 Resize 到 64 时的插值失真。这个技巧在图像分类里也常用。3. GAN 的数学原理把公式拆到能对应上每一行代码很多人写 GAN 代码时是照着抄的损失函数、优化器、更新顺序全靠背。这样做在标准配置下能跑通但一旦出现问题就完全不知道从哪下手。所以这一节我把公式和代码的对应关系讲透尤其是那个被搜索得最多的问题——原始 GAN 公式里的交叉熵为什么看不到负号。3.1 原始 GAN 的极小极大博弈GAN 的目标函数写成这样min_G max_D V(D, G) E_{x~p_data}[log D(x)] E_{z~p_z}[log(1 - D(G(z)))]拆开看判别器 D 要最大化这个值生成器 G 要最小化它。判别器希望D(x)接近 1对真图打出高分D(G(z))接近 0对假图打出低分这样两项的 log 都趋近于 0V 就大。生成器则希望D(G(z))接近 1这样第二项log(1-D(G(z)))趋近于负无穷V 就小。理论上这个博弈的纳什均衡点出现在p_g p_data此时判别器对任何输入都只能输出 0.5因为它无法区分。这也是为什么训练时你会看到判别器损失慢慢稳定在log 2 ≈ 0.693附近——那是个好信号说明两者接近势均力敌。3.2 交叉熵为什么没有负号这是被问得最多的一个问题答案其实很简单因为论文里写的是最大化代码里做的是最小化。二分类交叉熵的标准定义是BCE(p, y) -[y * log(p) (1-y) * log(1-p)]这里 y 是标签p 是模型预测为正类的概率。当 y1真样本时损失简化为-log(p)当 y0假样本时损失简化为-log(1-p)。现在把两者对照。判别器对真样本的目标是最大化log D(x)等价于最小化-log D(x)而-log D(x)正是 y1 时的交叉熵。判别器对假样本的目标是最大化log(1-D(G(z)))等价于最小化-log(1-D(G(z)))正是 y0 时的交叉熵。所以公式里的没有负号是因为它描述的是判别器奖励自己判对的方向而代码里的损失函数是惩罚自己判错的方向于是负号就显式地出现了。两者描述的是同一件事只是视角相反。代码里的体现是这样criterion nn.BCELoss() # ---- 判别器更新 ---- # 真样本标签为 1损失 -log(D(x)) real_pred D(real_img) loss_d_real criterion(real_pred, torch.ones_like(real_pred)) # 假样本标签为 0损失 -log(1 - D(G(z))) fake_pred D(fake_img.detach()) loss_d_fake criterion(fake_pred, torch.zeros_like(fake_pred)) loss_d loss_d_real loss_d_fake # ---- 生成器更新 ---- # 生成器希望判别器把假图判成真所以标签用 1 gen_pred D(fake_img) loss_g criterion(gen_pred, torch.ones_like(gen_pred))注意生成器的损失把假图的标签设成 1意思是我希望你把它判成真。这是原始论文里的非饱和形式比早期用的minimize log(1-D(G(z)))更实用——因为在训练初期 D 很容易把假图判得死死的log(1-D)的梯度会非常小非饱和形式则能提供更稳定的梯度。3.3 转置卷积的输出尺寸与感受野生成器从 128 维噪声一路放大到 64×64靠的是转置卷积。输出尺寸的公式是out (in - 1) * stride - 2 * padding kernel_size取stride2, padding1, kernel4代入得out (in-1)*2 - 2 4 2*in正好翻倍。所以从噪声的 1×1 出发每一层都翻倍经过 6 层就能到 641→4→8→16→32→64。注意第一层我用的 kernel4、stride1、padding0把 1×1 的噪声直接铺成 4×4和后续的翻倍模式错开这样才能凑出 4、8、16、32、64 这串数。判别器这边反过来用 stride2 的普通卷积逐层减半64→32→16→8→4→1。最后一层用 kernel4、stride1、padding0把 4×4 的特征图压成 1×1 的标量输出。整条路径下来最后一个卷积层的感受野刚好覆盖整个 64×64 输入这也是为什么输出能是一个有意义的全局判断。注意kernel、stride、padding 三者的组合不是随便凑的。如果输出尺寸除不尽网络会在拼接时直接报维度不匹配调试时间可能花掉一下午。建议先用纸笔算一遍每层尺寸再动手写代码。3.4 权重初始化与批归一化到底在解决什么DCGAN 论文里明确要求用正态分布初始化均值 0、标准差 0.02。为什么要专门做这件事因为 PyTorch 默认的卷积初始化是 Kaiming 均匀分布尺度和 0.02 的正态分布差距不小。在 GAN 这种没有明确监督信号、梯度又比较脆弱的场景里初始权重过大会让判别器在第一个 iteration 就输出极端置信度生成器接下来的梯度全部失效。def weights_init(m): name m.__class__.__name__ if name.find(Conv) ! -1: nn.init.normal_(m.weight.data, 0.0, 0.02) elif name.find(BatchNorm) ! -1: nn.init.normal_(m.weight.data, 1.0, 0.02) nn.init.constant_(m.bias.data, 0)批归一化的作用则更微妙。在生成器里它把每层的激活值重新标准化防止某些通道的数值越来越大最终让 Tanh 全部饱和到 ±1输出变成大色块。在判别器里它起到正则化的作用抑制模型对单个样本的过拟合。但有两处必须去掉生成器的输出层否则 Tanh 前的分布被强行拉平输出的颜色范围会失真和判别器的输入层否则批内所有样本被耦合在一起判别器会学到这批图整体长什么样而不是单张图的真假这是 GAN 训练里一个隐蔽的作弊路径。LeakyReLU 的负斜率取 0.2 也是经验值。ReLU 在负数区间梯度为 0判别器一旦某些神经元死掉就再也活不过来负斜率给了它们一条生路。4. 代码实现生成器、判别器与训练循环理论铺垫够了接下来是能直接跑的部分。整个项目的依赖很轻PyTorch、torchvision、Pillow、OpenCV、tqdm。除了 OpenCV 用于前期的数据清洗训练阶段只需要前两个。4.1 环境搭建我用 conda 管理环境因为 torch 的版本和 CUDA 驱动强绑定隔离环境能避免把系统里的其他项目搞乱。conda create -n animegan python3.10 -y conda activate animegan # CUDA 11.8 版本按自己的驱动选对应版本 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 # 纯 CPU 环境就用这一行 # pip install torch torchvision pip install opencv-python pillow tqdm tensorboard装完之后先跑一段验证确认版本和 GPU 可用import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0) if torch.cuda.is_available() else CPU only)如果cuda.is_available()返回 False先别急着怀疑代码。我遇到过三次两次是驱动版本比 wheel 要求的低一次是 conda 环境的 PATH 覆盖了系统 CUDA 路径。用nvidia-smi看驱动支持的 CUDA 版本上限再对照 PyTorch 官网的版本对照表选 wheel基本能定位。Windows 下还有个特有情况num_workers 0时 DataLoader 会走 spawn 模式如果数据加载逻辑写在if __name__ __main__外面会无限递归报错。这个报错的堆栈很长看着吓人实际加一行判断就好。4.2 生成器import torch.nn as nn class Generator(nn.Module): def __init__(self, nz128, ngf64, nc3): super().__init__() self.net nn.Sequential( # 输入: (nz, 1, 1) - (ngf*8, 4, 4) nn.ConvTranspose2d(nz, ngf * 8, 4, 1, 0, biasFalse), nn.BatchNorm2d(ngf * 8), nn.ReLU(True), # (ngf*8, 4, 4) - (ngf*4, 8, 8) nn.ConvTranspose2d(ngf * 8, ngf * 4, 4, 2, 1, biasFalse), nn.BatchNorm2d(ngf * 4), nn.ReLU(True), # - (ngf*2, 16, 16) nn.ConvTranspose2d(ngf * 4, ngf * 2, 4, 2, 1, biasFalse), nn.BatchNorm2d(ngf * 2), nn.ReLU(True), # - (ngf, 32, 32) nn.ConvTranspose2d(ngf * 2, ngf, 4, 2, 1, biasFalse), nn.BatchNorm2d(ngf), nn.ReLU(True), # - (nc, 64, 64) nn.ConvTranspose2d(ngf, nc, 4, 2, 1, biasFalse), nn.Tanh(), ) def forward(self, z): return self.net(z)有两个细节容易被忽略。第一所有转置卷积都设了biasFalse因为紧跟着就是批归一化偏置项会被 BN 的减均值操作完全抵消留着只是白白增加参数和显存占用。第二最后一层用 Tanh 而不是 Sigmoid前者的梯度在中心区域更强早期训练更稳。4.3 判别器class Discriminator(nn.Module): def __init__(self, nc3, ndf64): super().__init__() self.net nn.Sequential( # (nc, 64, 64) - (ndf, 32, 32) nn.Conv2d(nc, ndf, 4, 2, 1, biasFalse), nn.LeakyReLU(0.2, inplaceTrue), # - (ndf*2, 16, 16) nn.Conv2d(ndf, ndf * 2, 4, 2, 1, biasFalse), nn.BatchNorm2d(ndf * 2), nn.LeakyReLU(0.2, inplaceTrue), # - (ndf*4, 8, 8) nn.Conv2d(ndf * 2, ndf * 4, 4, 2, 1, biasFalse), nn.BatchNorm2d(ndf * 4), nn.LeakyReLU(0.2, inplaceTrue), # - (ndf*8, 4, 4) nn.Conv2d(ndf * 4, ndf * 8, 4, 2, 1, biasFalse), nn.BatchNorm2d(ndf * 8), nn.LeakyReLU(0.2, inplaceTrue), # - (1, 1, 1) nn.Conv2d(ndf * 8, 1, 4, 1, 0, biasFalse), nn.Sigmoid(), ) def forward(self, x): return self.net(x).view(-1)最后一层的.view(-1)把 (B,1,1,1) 展平成 (B,)因为 BCELoss 要求输入和目标的形状一致。这个形状问题是个高频报错点报错信息是 Target size must be the same as input size看到它就知道该加 view 了。判别器第一层不加 BN但加了 LeakyReLU这是 DCGAN 明文规定的。我实测过加 BN 的版本前几个 epoch 看起来损失更低但生成质量明显更差因为批内的真图和假图信息会被混合统计。4.4 训练循环import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import DataLoader from torchvision import transforms, utils from tqdm import tqdm import os device torch.device(cuda if torch.cuda.is_available() else cpu) NZ, BATCH, EPOCHS 128, 128, 120 LR, BETA1 2e-4, 0.5 transform transforms.Compose([ transforms.Resize(72), transforms.CenterCrop(64), transforms.RandomHorizontalFlip(0.5), transforms.ToTensor(), transforms.Normalize((0.5,) * 3, (0.5,) * 3), ]) dataset AnimeFaceDataset(data/faces, transformtransform) loader DataLoader(dataset, batch_sizeBATCH, shuffleTrue, num_workers4, pin_memoryTrue, drop_lastTrue) G Generator(NZ).to(device) D Discriminator().to(device) G.apply(weights_init) D.apply(weights_init) criterion nn.BCELoss() opt_g optim.Adam(G.parameters(), lrLR, betas(BETA1, 0.999)) opt_d optim.Adam(D.parameters(), lrLR, betas(BETA1, 0.999)) fixed_z torch.randn(64, NZ, 1, 1, devicedevice) # 固定噪声用来看演变 os.makedirs(samples, exist_okTrue) os.makedirs(ckpt, exist_okTrue) for epoch in range(EPOCHS): for i, real in enumerate(tqdm(loader, descfEpoch {epoch1}/{EPOCHS})): real real.to(device, non_blockingTrue) bs real.size(0) # ---------- 更新判别器 ---------- D.zero_grad(set_to_noneTrue) real_label torch.full((bs,), 0.9, devicedevice) # 标签平滑 fake_label torch.full((bs,), 0.1, devicedevice) out_real D(real) loss_d_real criterion(out_real, real_label) z torch.randn(bs, NZ, 1, 1, devicedevice) fake G(z) out_fake D(fake.detach()) loss_d_fake criterion(out_fake, fake_label) loss_d loss_d_real loss_d_fake loss_d.backward() opt_d.step() # ---------- 更新生成器 ---------- G.zero_grad(set_to_noneTrue) out_gen D(fake) loss_g criterion(out_gen, torch.ones(bs, devicedevice)) loss_g.backward() opt_g.step() # 每轮存一张样本网格 if (epoch 1) % 5 0 or epoch 0: G.eval() with torch.no_grad(): grid utils.make_grid(G(fixed_z), nrow8, normalizeTrue) utils.save_image(grid, fsamples/epoch_{epoch1:03d}.png) G.train() torch.save({g: G.state_dict(), d: D.state_dict(), epoch: epoch}, fckpt/ckpt_{epoch1:03d}.pth)几个关键点解释一下。标签平滑真标签用 0.9 而不是 1.0假标签用 0.1 而不是 0.0。这是防止判别器变得过度自信的最简单手段。当判别器对真图输出 0.999 时-log(0.999)几乎为 0梯度也就几乎消失用 0.9 做目标能让它一直保持一点不确定梯度不至于消失。实测这一招对训练稳定性的提升很明显几乎是必开的。betas(0.5, 0.999)Adam 默认的 beta1 是 0.9但在 GAN 里会带来问题——动量的累积让参数更新有滞后性判别器和生成器之间的动态博弈需要一个更灵敏的响应所以调到 0.5。这是 DCGAN 论文给出的设置至今仍然是 GAN 训练的基线配置。fake.detach()更新判别器时必须切断假图到生成器的梯度否则一次 backward 会同时更新两个网络的参数训练会彻底失控。很多人第一次写 GAN 就是漏了这个 detach表现为损失曲线看起来正常但生成质量原地踏步。更新顺序先判别器后生成器且都是每个 iteration 更新一次。这个 1:1 的比例在 DCGAN 上工作良好但如果你发现判别器明显学得太快可以改成判别器每更新一次、生成器更新两次。4.5 采样、保存与断点续训固定噪声向量这一段是排查问题的关键手段。每次采样都用同一批 z就能通过样本图的变化直观看到训练轨迹。正常情况下你会看到前 5 轮是一片混沌的彩色噪声第 10 到 20 轮开始出现明显的色块分区比如上面深色下面浅色对应头发和皮肤第 30 轮之后轮廓逐渐清晰第 60 轮左右能看出五官。如果从第 20 轮开始连续 5 轮采样图完全没有变化那基本可以判定是模式崩溃了不用再等着看。断点续训的加载逻辑ckpt torch.load(ckpt/ckpt_060.pth, map_locationdevice) G.load_state_dict(ckpt[g]) D.load_state_dict(ckpt[d]) start_epoch ckpt[epoch]顺便说一句优化器的状态最好也一起存。不过实测下来Adam 的动量状态在中断后重新累积对最终效果的影响很有限为了省显存我一般只存模型权重。5. 训练调优与问题排查实录跑过几十轮实验之后你会发现 GAN 的问题基本都是那几类而且都有相对固定的信号可以识别。这一节按症状分类把我实际遇到过的情况整理出来。5.1 模式崩溃的识别与应对模式崩溃的表现很好认生成器输出的所有样本高度相似可能全是同一个发色、同一个构图甚至背景颜色都一样。本质原因是生成器发现了一个能稳定骗过当前判别器的捷径于是把所有噪声输入都映射到这一个点上放弃了探索整个分布。从损失上看模式崩溃时通常伴随判别器损失快速下降并稳定在一个很低的值比如 0.1 以下生成器损失则持续上升。这很好理解——生成器只会做那一张图判别器很快学会了识破但因为生成器不再变化判别器也不会再遇到新的挑战。我的处理顺序是这样的第一步先降低判别器的学习率比如从 2e-4 降到 1e-4看能不能把平衡拉回来第二步加标签平滑如果之前没加的话第三步在判别器输入上加一点高斯噪声标准差从 0.05 开始试这个做法相当于人为模糊真假边界让判别器没法记住具体样本第四步如果还没改善就说明判别器容量相对于数据量太大了把 ndf 从 64 降到 32。5.2 判别器过强或过弱的判断判断这个问题的依据是损失曲线的形态我整理了一张对照表现象生成器损失判别器损失判断处理方式正常训练1.5 到 4.0 之间波动0.4 到 0.9 之间波动势均力敌保持现状判别器过强持续上升超过 6低于 0.1 并稳定生成器拿不到梯度降低 D 学习率、加标签平滑、给 D 减容量判别器过弱持续下降低于 0.8持续在 1.2 以上生成器在骗自己提高 D 学习率、给 G 减容量模式崩溃缓慢上升很低但不再下降生成器退化见 5.1 的处理顺序有个容易被误判的情况训练前 10 轮判别器损失很低是正常的因为这个阶段生成器还没学会生成任何有意义的结构判别器区分真假确实很容易。不要在这个阶段就急着调参至少等到 20 轮之后再判断趋势。另一个经验是观察生成器损失的时候不要只看绝对值要看它的波动幅度。健康状态下的生成器损失是持续抖动的因为每个 batch 的假图都在变如果它稳定在一个数值上几乎不动多半是梯度已经断了。5.3 报错速查表下面这些报错我基本都踩过按出现频率排序报错信息原因解决方式CUDA out of memorybatch size 或通道数过大降到 64 或 32开启torch.cuda.empty_cache()Target size must be the same as input size判别器输出形状没展平输出后加.view(-1)Given groups1, weight of size...通道数在层间不匹配检查 ngf/ndf 的乘法是否对齐DataLoader worker exited unexpectedlyWindows 下多进程问题数据加载放进__main__或 num_workers 设为 0loss 变成 nan学习率过大或对数运算遇到 0降学习率损失用BCEWithLogitsLoss更稳Input type (torch.cuda.FloatTensor) ...模型和数据不在同一设备检查to(device)有没有漏关于 nan还有一个隐蔽来源如果判别器在某一批上对假图输出了精确的 0 或 1log(0)会产生负无穷。虽然 PyTorch 的 BCELoss 内部做了 clamp但在混合精度训练下这个 clamp 有可能失效。我的做法是遇到 nan 就暂时关掉混合精度或者改用BCEWithLogitsLoss并把判别器最后一层的 Sigmoid 去掉数值稳定性会好很多。提示混合精度训练确实能把显存占用降到六成、速度提升三成左右但在 GAN 上要格外小心。判别器的损失在真假之间跨度很大FP16 的表示范围容易在早期就溢出。建议先用全精度跑通一个完整流程确认效果好再尝试 fp16。6. 效果再往上走从 DCGAN 到更稳的方案DCGAN 跑通之后如果你想让生成质量再上一个台阶有几条路可以走。我按投入产出比从高到低排一下。6.1 WGAN-GP解决训练不稳定最有效的方案WGAN 的核心改动是把 JS 散度换成 Wasserstein 距离判别器最后一层去掉 Sigmoid此时它不叫判别器而叫 Critic输出一个实数分数。配合梯度惩罚项训练稳定性会有质的提升模式崩溃的概率大幅降低。关键改动只有三处def gradient_penalty(D, real, fake, device, lambda_gp10.0): alpha torch.rand(real.size(0), 1, 1, 1, devicedevice) interpolated (alpha * real (1 - alpha) * fake).requires_grad_(True) score D(interpolated) grad torch.autograd.grad( outputsscore, inputsinterpolated, grad_outputstorch.ones_like(score), create_graphTrue, retain_graphTrue, only_inputsTrue, )[0] grad grad.view(grad.size(0), -1) gp ((grad.norm(2, dim1) - 1) ** 2).mean() * lambda_gp return gp第一处是判别器去掉 Sigmoid第二处是损失函数改成loss_d fake_score.mean() - real_score.mean() gp注意这里是 Critic 要最小化它相当于最大化真分减假分第三处是优化器换成 RMSprop 或 Adam 但学习率降到 1e-4同时把 betas 调成 (0.5, 0.9)。我在同样的数据集上对比过DCGAN 在第 40 轮左右会出现一次明显的质量回落然后重新爬起来WGAN-GP 全程曲线平滑得多。代价是每个 iteration 要多算一次梯度惩罚训练时间大概增加 30%。6.2 条件生成与风格控制无条件 GAN 的问题是噪声到图像的映射完全不可控你想要金发的结果可能出来一片蓝发。解决办法是引入条件信息也就是 cGAN。最常用的做法是标签嵌入后与噪声拼接或者在判别器里用过 Projection 判别器把标签信息以点积形式注入。如果你不想额外标注数据还有一个取巧的做法用现成的图像特征提取器给每张训练图算一个聚类标签比如按发色分成 6 类然后拿这个标签做条件训练。这样生成时就能指定发色实用性提升很大。6.3 怎么客观评估生成质量靠肉眼看样本图只是最粗的评估方式。有几个指标相对成熟FID 衡量生成分布和真实分布在特征空间的距离数值越低越好一般二次元头像数据集上能降到 30 以内就算不错IS 衡量生成样本的清晰度和多样性但它在人脸类数据上参考价值有限因为人脸本身的类别区分度不高。实际项目里我一般两种手段并用每 10 轮算一次 FID 画成曲线同时保留固定噪声的采样图。FID 告诉你整体趋势采样图告诉你具体在哪些方面退化了——比如 FID 在下降但采样图里的眼睛开始糊那就说明模型把精力都花在了头发和背景上。算 FID 的时候有个坑特征提取器通常是在真实照片上训练的直接拿来算二次元头像的 FID绝对值会偏高。这时候更重要的是看相对趋势而不是绝对数值同一套特征提取器下不同 epoch 之间的比较才有意义。关于后续还能怎么扩展我个人的做法是先把 DCGAN 的基线固定下来把它当做一个基准之后每次改动只动一个变量并且用同样的随机种子跑对比。这个习惯看着笨但能让你真正搞清楚每个改动到底带来了什么而不是凭感觉以为某个技巧有用。踩过几次改了三处效果变好了但不知道是哪一处的功劳的坑之后我现在连换优化器都要单独跑一组对照。训练 GAN 这件事耐心和纪律比技巧更重要。
返回列表