十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

用GAN做图像增强:Kaggle竞赛中的数据增强实战指南

用GAN做图像增强:Kaggle竞赛中的数据增强实战指南 简介面向 Kaggle 图像分类与生成任务的项目资源聚焦 FER13 数据集 7 类情感识别中样本不均衡与增强不足的问题。作者提出通过 GAN 生成更多图像并进行类均衡从而提升简单 CNN 模型在整体测试集上的多分类准确率适合想了解 DCGAN 数据增强或复现情感分类实验的机器学习初学者与进阶者。压缩包共 19 个文件含 7 个 Python 脚本如 dcgan.py、training.py、evaluation.py、make_csv.py、resize_images.py、8 张 PNG 结果图、1 张 JPG 图片以及 README 与 requirements.txt整体仅 148KB结构紧凑。已有 4203 人学习/下载。资源提供可直接运行的训练与评估流程从图像缩放、CSV 制作、DCGAN 生成到模型训练与准确率评估均有对应脚本README 可快速上手PNG 结果图便于对比生成样本和原始情绪样本。适合作为课程设计、Kaggle 入门或情感识别方向的基础参考也可基于这份代码扩展网络结构或更换数据集快速验证 GAN 增强对分类性能的影响。1. 用 GAN 做图像增强本质上是在补分布而不是改像素很多 Kaggle 项目并不是在刷 model architecture而是在处理数据一侧的失衡。比如正样本只有 200 张翻转、染色、加噪这些传统增强能增加数量却不能增加纹理变化模型很快把背景当特征。用 GAN 做图像增强核心不是让生成器随便画图而是让它学真实样本的纹理、光照与边缘分布再造出新样本——这就是“图像生成”和“图像增强”在 GAN 这里交会的原因。不过 GAN 的训练不稳定参数设错会导致生成图像千篇一律或直接崩成噪点。下面按原理、选型、训练、评估、推理五段展开给出一条在 Kaggle 项目里能落地的路径。2. GAN 图像增强原理与选型先判断问题类型再选模型2.1 对抗训练的本质是分布逼近GAN 由一个生成器 G 和一个判别器 D 组成。G 接收一个随机向量 z输出一张图像 G(z)D 接收一张图像输出它来自真实分布的概率。训练时 D 尽量区分真实样本和生成样本G 则尽量骗过 D。两边都在优化同一个极小极大目标最终 G 生成的图像会落到真实图像的流形上。在 Kaggle 增强场景里这个“流形”通常比我们手里的几千张样本更光滑。这里要注意一点如果用纯随机噪声生成图像你无法控制生成的是正样本还是负样本。所以在数据增强任务里无条件 GAN 的实用性受限最常用的是条件 GANcGAN。条件 GAN 把某些已知信息作为输入比如一张模糊图、一张夜间图、一个分割 mask然后生成对应的清晰图、日间图、原图。这相当于让生成器在给定条件下做“增强”而不是无中生有。2.2 用表格选型DCGAN、pix2pix、CycleGAN、SRGAN对抗生成网络 GAN 系列在图像增强上的分支主要是 pix2pix、CycleGAN 和 SRGAN 三个方向。在 Kaggle 项目里通常先看自己的数据是否成对。比如你有低光图和正常光图就是成对的直接上 pix2pix只有一堆夜间图想变成白天图没有配对目标用 CycleGAN想把小目标放大再送去检测用 SRGAN 或 ESRGAN。我把常用模型放在一张表里模型系列输入条件典型任务数据配对要求Kaggle 适用场景DCGAN纯噪声无条件生成不需要罕见样本扩充但类别不受控pix2pixcGAN一张待处理图去噪、去模糊、补光需要严格配对低光增强、去噪后做分类CycleGAN一张待处理图风格迁移两个域无需配对把模拟数据转成真实风格SRGAN / ESRGAN低分辨率图超分辨率需要 LR/HR 配对小目标检测前的预处理这个表格基本决定了你的项目走向。我的建议是能凑出配对数据就用 pix2pix凑不出来先试试传统的直方图均衡化或小波变换增强不要一上来就 CycleGAN。CycleGAN 的循环一致性损失会明显拉长训练时间在 Kaggle 的 kernel 限时下经常跑不完。图像增强算法有很多种GAN 只是其中一种思路选型时要先算清时间预算。2.3 对抗损失与重建损失要做加法只用对抗损失训练生成图在颜色和纹理上可以很逼真但结构上与输入条件没有强绑定。比如做去模糊增强时生成器可能甩掉眉毛、牙齿等局部细节虽然判别器看来像人脸实际对分类器不利。为了保持输入输出的像素级对应pix2pix 给生成器加了一个 L1 重建损失通常权重设成 100。这一步是 Kaggle 项目里少踩一半坑的关键。2.3.1 损失函数组合的最小示例import torch import torch.nn as nn l1_loss nn.L1Loss() gan_criterion nn.BCEWithLogitsLoss() # 生成器输入 low 图目标 true_img l1_term l1_loss(fake_img, true_img) gan_term gan_criterion(d_fake_output, torch.ones_like(d_fake_output)) g_loss gan_term 100 * l1_term这段代码里gan_term是让生成器学会骗过判别器的对抗部分l1_term是强制生成结果在像素级上靠近目标。系数100是 pix2pix 论文里的经典默认值在多数数据集上可以直接用。如果发现生成图像边缘太糊可以把它降到 50 或 80如果图像颜色纹理很怪而轮廓很准就适当加大到 150。2.3.2 为什么 L1 权重常设成 100判别器的输出是一个概率图每个位置对应原图的一个 patch。当 patch 数目很多时整个对抗损失的量级会被放大所以不能简单地把两个损失按 1:1 相加。这也是为什么 L1 权重必须设成一个大数或者反过来调小 GAN 损失。你可以在训练循环里记录两项损失的均值观察它们是否在同一个数量级。如果差距过大调整gan_term的缩放而不是盲目改学习率。3. 在 Kaggle Notebook 跑通一条最小 GAN 图像增强流水线3.1 用 Dataset 管理 Kaggle 数据集的配对关系Kaggle 数据集大多以文件夹形式给到。你要先确认文件名能否一一对应比如low/0001.png对应high/0001.png。下面这个 Dataset 会把两组图像装进同一个列表在__getitem__里同时返回低质量与高质量图像。from torch.utils.data import Dataset from PIL import Image import glob class PairedImageDataset(Dataset): def __init__(self, low_dir, high_dir, transformNone): self.low_paths sorted(glob.glob(low_dir /*.png)) self.high_paths sorted(glob.glob(high_dir /*.png)) assert len(self.low_paths) len(self.high_paths), ( f配对数量不一致: {len(self.low_paths)} vs {len(self.high_paths)} ) self.transform transform def __len__(self): return len(self.low_paths) def __getitem__(self, idx): low Image.open(self.low_paths[idx]).convert(RGB) high Image.open(self.high_paths[idx]).convert(RGB) if self.transform: low, high self.transform(low, high) return low, high这里有三个可调点文件扩展名、通道数和 transform。Kaggle 上很多医学或遥感图是灰度图.convert(RGB)会把单通道复制成三通道省去改网络输入通道的麻烦。transform 里要做Resize和ToTensor注意两张图必须使用同一个随机种子做水平翻转否则低图和目标图会错位。3.2 生成器结构用 U-Net 保边缘还是用 ResNet 省显存pix2pix 生成器普遍用 U-Net因为跳跃连接能把输入的低级细节直接送到解码器生成结果更锐利。但 U-Net 的显存开销比 ResNet 生成器高Kaggle 的免费 GPU 跑 256x256 图像时batch size 往往只能设 1。我一般会写一个轻量 U-Net让每个 block 的通道数不要超过 512。import torch.nn as nn class Generator(nn.Module): def __init__(self, in_channels3, out_channels3): super().__init__() # 下采样 self.down1 nn.Sequential(nn.Conv2d(in_channels, 64, 4, 2, 1), nn.LeakyReLU(0.2)) self.down2 nn.Sequential(nn.Conv2d(64, 128, 4, 2, 1), nn.InstanceNorm2d(128), nn.LeakyReLU(0.2)) self.down3 nn.Sequential(nn.Conv2d(128, 256, 4, 2, 1), nn.InstanceNorm2d(256), nn.LeakyReLU(0.2)) # 上采样 self.up1 nn.Sequential(nn.ConvTranspose2d(256, 128, 4, 2, 1), nn.InstanceNorm2d(128), nn.ReLU()) self.up2 nn.Sequential(nn.ConvTranspose2d(256, 64, 4, 2, 1), nn.InstanceNorm2d(64), nn.ReLU()) self.up3 nn.Sequential(nn.ConvTranspose2d(128, out_channels, 4, 2, 1), nn.Tanh()) def forward(self, x): d1 self.down1(x) d2 self.down2(d1) d3 self.down3(d2) u1 self.up1(d3) u2 self.up2(torch.cat([u1, d2], dim1)) u3 self.up3(torch.cat([u2, d1], dim1)) return u3这个生成器非常小适合先跑通流程。注意up2和up3都做了cat这就是 U-Net 的跳跃连接。下采样用InstanceNorm2d而不是BatchNorm2d因为 batch size 通常是 1 或 2BN 统计均值偏差大。如果显存还有富余可以把up3输出前再接一个 3x3 卷积来减少棋盘伪影。3.3 PatchGAN 判别器的原理与代码判别器不需要看整张图只要输出一个 NxN 的概率矩阵。每一个输出值都对应原图一个 patch网络越深patch 的感受野越大。PatchGAN 的好处是参数少、训练快对高频纹理敏感很适合增强任务。class Discriminator(nn.Module): def __init__(self, in_channels6): super().__init__() # 输入 low 图和 fake/real 图拼接共 6 通道 self.model nn.Sequential( nn.Conv2d(in_channels, 64, 4, 2, 1), nn.LeakyReLU(0.2), nn.Conv2d(64, 128, 4, 2, 1), nn.InstanceNorm2d(128), nn.LeakyReLU(0.2), nn.Conv2d(128, 256, 4, 2, 1), nn.InstanceNorm2d(256), nn.LeakyReLU(0.2), nn.Conv2d(256, 1, 4, 1, 1), ) def forward(self, low_img, target_img): return self.model(torch.cat([low_img, target_img], dim1))这里的核心参数是in_channels6把条件图和目标图按通道拼接。判别器的最后一层没有 Sigmoid因为后面要配合BCEWithLogitsLoss。如果 Kaggle 训练时发现判别器过度强于生成器可以在每个下采样层后加 Dropout2d或把通道数从 64 改成 32。3.4 训练循环与参数表训练时每步先更新判别器再更新生成器。更新判别器时要把生成图从计算图中分离更新生成器时则要保留生成图与判别器之间的梯度。下面是最小循环片段for epoch in range(epochs): for low, high in train_loader: low, high low.cuda(), high.cuda() # 生成假图 fake generator(low) # 判别器训练 d_real discriminator(low, high) d_fake discriminator(low, fake.detach()) d_loss gan_criterion(d_real, torch.ones_like(d_real)) \ gan_criterion(d_fake, torch.zeros_like(d_fake)) d_optimizer.zero_grad() d_loss.backward() d_optimizer.step() # 生成器训练 d_fake2 discriminator(low, fake) g_gan gan_criterion(d_fake2, torch.ones_like(d_fake2)) g_l1 l1_loss(fake, high) g_loss g_gan 100.0 * g_l1 g_optimizer.zero_grad() g_loss.backward() g_optimizer.step()注意生成器训练里没有fake.detach()这是刻意保留梯度。判别器训练里的fake.detach()是为了让判别器的反向传播不经过生成器二者不要混。Adam 优化器的betas(0.5, 0.999)是生成对抗网络里的常见设置标准 Adam 的0.9在 GAN 里更容易导致判别器反超。Kaggle 环境下的推荐参数参数推荐值备注图像尺寸256x256512 需要更大显存和更长训练时间batch size1 或 2PatchGAN 对小 batch 不敏感学习率2e-4按 discriminator 收敛快慢微调Adam beta10.5降低动量的历史累积L1 lambda100先用默认再按清晰度调整训练轮数30 起通常 30 轮能看到明显效果100 轮收敛提示先用epochs1跑一遍循环确认 forward 和 backward 维度没问题再把轮数调上去生成器在训练早期容易全输出均值图。3.5 把生成结果缓存下来避免重复训练生成器训练好之后用torch.no_grad()批量跑一遍训练集和验证集把生成图像以.npy或.png写到/kaggle/working/下。后续训练下游分类器时直接读取缓存不再跑生成器。这样可以让训练时间缩短一个数量级。generator.eval() with torch.no_grad(): for i, (low, _) in enumerate(all_loader): low low.cuda() fake generator(low).cpu() # 保存为 png保持文件名与原图一致 torch.save(fake, fcache/{i:05d}.pt)这段代码要注意all_loader的 shuffle 必须设为 False并且 batch 序号要和原文件名一致否则后面和标签对齐时会错位。缓存格式用.pt或.npy都比 PNG 快但可视化和调试不如 PNG 直观。4. 图像增强效果评估与调优FID、混合比例和模式崩溃4.1 用 FID 判断“生成分布是否接近真实分布”生成图像光“看着像”还不够下游分类器需要生成样本与真实样本在特征空间上同分布。FIDFréchet Inception Distance用预训练 Inception 提取特征计算两组特征均值与协方差的距离。FID 越低说明生成分布越接近真实分布在 Kaggle 的增强场景里比 IS 更适合因为 IS 依赖类别而我们在做回归增强时没有类别。from torchmetrics.image.fid import FrechetInceptionDistance fid FrechetInceptionDistance(feature2048) # 每次喂数据时先归一化到 [0,1] 再传入 fid.update(real_images, realTrue) fid.update(fake_images, realFalse) print(fFID {fid.compute().item():.2f})这里的feature参数控制 Inception 特征维度一般是 2048。更新时真实与生成样本数量要尽量接近否则协方差估计会偏移。FID 计算需要相同尺寸的输入如果你的图不是 299x299记得在更新前做 resize。注意 FID 只是分布相似度指标不直接代表下游准确率所以最终还是要用增强后的数据跑一版 CV。也可以在超分辨率任务中结合 PSNR 与 SSIMPSNR 看重像素误差SSIM 看重结构相似性。生成图像如果边缘锐利但过曝PSNR 会下降SSIM 却能捕捉到结构保真度所以两个指标要一起看。4.2 生成样本与原始样本的比例不是越多越好很多人拿到生成图后一股脑全加入训练集最后发现 CV 分数反而下降。原因很可能是生成样本大多集中在真实样本周围的“安全区”对分类器没有新信息。我一般做法是控制增强后总样本数不变先用 10% 生成样本替换掉重复性最高的真实样本再逐步上调。下面是一个单次三折交叉验证的示意结果配置验证集准确率备注原始数据 传统翻转裁剪0.812基线加入 15% GAN 增强样本0.833收益明显加入 30% GAN 增强样本0.829已经开始噪声过拟合加入 50% GAN 增强样本0.804生成样本重复度过高这张表说明 15% 附近往往是甜点。具体比例依赖任务复杂度和生成图像多样性我建议按 5% 步进做一次小网格搜索而不是拍脑袋定 30%。4.3 发现模式崩溃什么样的崩坏信号要立刻停训练 GAN 经常遇到三种信号生成图全部是同一张脸或同一纹理判别器 loss 快速掉到 0以及生成图像出现彩色斑块。第一种是模式崩溃第二种是判别器太强第三种一般是归一化错误或 tanh 输出与数据分布不匹配。遇到这些情况不要急着调网络结构先做三件事把学习率降到 1e-4 或改用betas(0.5, 0.999)。检查数据对的文件名顺序是否错位排序算法不一致会造成配对错乱特征学不到。每隔 20 步保存生成图像回放时序而不是看单张。如果想在损失函数层面缓解模式崩溃可以把判别器改成 WGAN-GP 的梯度惩罚代码改动不大def gradient_penalty(discriminator, low_img, real_img, fake_img): epsilon torch.rand(real_img.size(0), 1, 1, 1, devicereal_img.device) interpolated epsilon * real_img (1 - epsilon) * fake_img pred discriminator(low_img, interpolated) grad torch.autograd.grad(outputspred, inputsinterpolated, grad_outputstorch.ones_like(pred), create_graphTrue)[0] return ((grad.norm(2, dim1) - 1) ** 2).mean()这段梯度惩罚在判别器损失后面乘上一个系数lambda_gp常见是 10。注意 WGAN-GP 的对抗损失不再用 BCE而是直接让判别器输出一个实数距离。Kaggle 环境下这套改动会慢一些但能明显降低调参的痛苦。5. 进阶技巧把 GAN 增强缓存与置信度过滤组合起来5.1 训练时不要实时跑生成器预先离线生成很多新手把生成器直接放进训练循环每一步都 forward 一次。这在固定 epoch 的小数据集上还能跑但 Kaggle 的 GPU 时间配额有限生成器 forward 一次就消耗一次前向计算。正确做法是训练完生成器后离线批量增强并缓存成.pt文件之后再训练分类器时只做读取和常规增强。class CachedGANAugment: def __init__(self, cache_path): self.cache torch.load(cache_path, weights_onlyFalse) def __call__(self, image_tensor): # image_tensor 来自训练集cache 已经按同样顺序生成 aug self.cache[len(self.cache) % len(self.cache)] return torch.clamp(image_tensor * 0.7 aug * 0.3, 0, 1)这里的融合方式只是一种示例你可以按概率直接选择替换也可以做 alpha 混合。关键点是缓存到本地后下游训练循环里不需要再 import 生成器代码量小且不容易发生内存泄漏。5.2 用判别器置信度过滤掉“太假”的生成样本生成器并不是越训练越完美偶尔会输出边缘混沌的图像。与其人工看几百张图不如用训练好的判别器挑样本对每一张生成图用同一个判别器输出概率图统计平均置信度然后过滤掉低于某个阈值的样本。这个技巧能直接提升下游分类器的鲁棒性。def filter_by_discriminator(generator, discriminator, loader, threshold0.7): generator.eval(); discriminator.eval() selected_idx [] with torch.no_grad(): for idx, (low, high) in enumerate(loader): low low.cuda() fake generator(low) pred torch.sigmoid(discriminator(low, fake)) score pred.mean().item() if score threshold: selected_idx.append(idx) return selected_idx阈值threshold我一般取 0.7-0.8。如果选的样本太少调低到 0.5如果样本太多但质量差调高到 0.9。这个流程的本质是用判别器当质量评估器将“生成器有没有骗过我”转化为“这张图值不值得参与训练”在很多 Kaggle 图像竞赛里能省下大量人工筛选时间。你可以把过滤后的索引传回数据集再跑一轮普通训练循环此时训练时间和内存占用都远小于 GPU 实时增强方案。本文还有配套的精品资源点击获取
返回列表