ARTICLE DETAIL

资讯详情

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

DCGAN图像修复实战:本科生可复现的毕设级PyTorch项目

DCGAN图像修复实战:本科生可复现的毕设级PyTorch项目 简介本资源是一套基于Python实现的深度生成对抗网络GAN图像修复模型完整项目专为计算机相关专业本科生毕业设计、期末大作业及AI实战学习者打造。项目聚焦图像破损区域智能重建任务涵盖DCGAN架构搭建、生成器与判别器训练、数据预处理及效果可视化等核心环节难度适中且经过助教审定适合初涉生成式AI的学生快速上手与二次开发。压缩包共7个文件含6个Python源码如model.py定义网络结构、train-dcgan.py实现训练流程、complete.py提供推理接口和1份Markdown文档README.md含环境配置、运行说明与结果示例总大小仅12KB轻量易部署。目前已有164人学习下载所有代码均经本地编译验证可直接运行附带清晰模块划分与关键注释显著降低调试门槛助力高效完成课程实践与项目交付。1. 这不是玩具GAN一个能跑通、能改、能交差的图像修复实战项目专治毕设卡在“训练不收敛”和“eval报错no module”你是不是也试过从GitHub clone一个标着“GAN图像修复”的项目pip install完依赖一跑train-dcgan.py就卡在RuntimeError: Expected 4-dimensional input, but got 3-dimensional input或者训了200轮生成图全是灰色噪点loss曲线像心电图一样乱跳别急——这个基于Python实现的深度生成对抗网络GAN图像修复模型不是那种“README写得天花乱坠、代码里藏着三处硬编码路径”的半成品。它是一份经导师签字确认、评审98分、本地实测可复现的完整交付物从数据预处理complete.py、DCGAN核心结构model.py、梯度操作封装ops.py到训练主循环train-dcgan.py全部模块化、参数可调、日志可查。它不追求SOTA指标但严格遵循课程设计边界——用PyTorchNumPy实现不依赖TensorFlow或Keras所有tensor shape都显式校验batch_size、lr、noise_dim等关键参数全在train-dcgan.py顶部集中配置。适合计算机/人工智能方向本科生做毕业设计、期末大作业也适合想亲手拆解GAN训练黑匣子的初学者。你不需要懂Wasserstein距离但得会看loss下降趋势你不用重写判别器但能快速替换为ResNet骨干你甚至可以只跑simple-distributions.py验证高斯噪声生成逻辑——它就是那种“打开就能跑、跑完能讲清、答辩能答住”的务实型源码包。2. 从零启动环境准备、目录结构解析与核心模块职责拆解2.1 环境搭建避开CUDA版本陷阱的Python依赖清单这个项目对环境要求明确且克制Python 3.7–3.9不兼容3.10因ops.py中部分torch.nn.functional调用在新版有行为变更PyTorch 1.8.1cu111必须带CUDA支持CPU版训DCGAN极慢且易OOM。我建议用conda新建隔离环境conda create -n gan-repair python3.8 conda activate gan-repair pip install torch1.8.1cu111 torchvision0.9.1cu111 -f https://download.pytorch.org/whl/torch_stable.html pip install numpy opencv-python tqdm matplotlib scikit-image提示不要用pip install torch自动匹配最新版项目中model.py的nn.ConvTranspose2dstride设置依赖1.8.1的padding计算逻辑新版会引发output size mismatch。若你只有CPU需手动修改train-dcgan.py第32行device torch.device(cpu)并把batch_size从64降到16否则simple-distributions.py里的torch.randn(64, 100, 1, 1)会爆内存。2.2 目录结构即设计蓝图每个文件解决什么问题整个source.zip解压后是扁平结构但模块职责清晰绝非脚本堆砌文件名核心职责关键技术点是否可独立运行utils.py数据加载与增强读取图像、裁剪为64×64、归一化到[-1,1]、添加随机mask模拟破损使用cv2.resize双线性插值 np.random.rand生成mask概率图✅ 可单独测试load_data()输出shapemodel.pyDCGAN生成器G与判别器D定义G用5层ConvTranspose2d上采样D用5层Conv2d下采样均含BatchNorm和LeakyReLU所有卷积层stride2保证尺寸翻倍/减半padding1避免边缘失真✅python model.py会打印G/D结构ops.py训练原子操作封装gradient_penaltyWGAN-GP用、get_optimizerAdam with betas(0.5,0.999)、save_checkpoint保存G/D state_dictepochlosstorch.autograd.grad手动求二阶导torch.save含torch.version.cuda校验❌ 依赖model.py但函数可单元测试train-dcgan.py主训练循环加载数据、初始化G/D、交替训练1步D1步G、每10轮保存checkpoint、每50轮生成sample图torch.no_grad()包裹G生成、torch.set_grad_enabled(True)控制D梯度✅ 设置--epochs 10即可快速验证流程simple-distributions.py验证生成器基础能力用标准正态分布z生成fake image不接判别器纯看G能否把噪声映射成结构化图像torch.randn(64,100,1,1)→ G →torch.sigmoid→cv2.imwrite✅ 首推运行此文件5秒出图确认G无bug2.3 模块间数据流一张图看懂tensor如何穿越GAN训练时tensor流动严格遵循DCGAN范式utils.py输出(64,3,64,64)张量batch64RGB64×64model.py中G接收(64,100,1,1)噪声向量输出(64,3,64,64)fake_imgD接收real_img或fake_img输出(64,1)logits未sigmoidops.py中get_optimizer为G/D分别创建Adam实例gradient_penalty仅在D更新时注入WGAN-GP约束关键细节train-dcgan.py第87行fake_img G(noise).detach()中的.detach()切断G梯度确保D训练时不更新G参数——这是GAN训练稳定性的基石。若此处漏掉D loss会异常震荡这是新手最常踩的坑之一。3. 训练全流程实操从数据准备到loss可视化每一步命令都附参数说明3.1 数据准备三类输入路径适配与mask生成逻辑项目默认读取./data/下的图像但utils.py支持三种模式模式1单图修复调试用将一张test.jpg放入./data/load_data()自动resize为64×64mask生成逻辑在add_mask()函数def add_mask(img): # img shape: (3,64,64) mask np.random.rand(64,64) 0.7 # 30%像素置0 masked_img img * mask[None,:,:] # 广播到3通道 return masked_img, mask参数说明0.7是mask保留率值越小破损越严重mask[None,:,:]增加channel维度适配RGB。模式2批量修复毕设正式用在train-dcgan.py第25行修改data_path ./data/train/要求该目录下全是.jpg/.png数量≥200张太少会导致D过拟合。模式3自定义mask进阶需求若你有特定破损模板如划痕、文字遮挡可替换add_mask()为mask cv2.imread(./masks/scratch.png, cv2.IMREAD_GRAYSCALE) # 64x64 binary masked_img img * (mask 128)[None,:,:]3.2 启动训练命令行参数详解与典型配置组合train-dcgan.py支持以下关键参数全部有默认值但建议显式指定参数默认值推荐值作用说明--epochs200100训练总轮数毕设100轮足够观察收敛趋势--batch_size6432显存不足时必调32对应约4GB显存--lr0.00020.0001G/D学习率过大导致loss爆炸过小收敛慢--beta10.50.5Adam beta1保持0.5是DCGAN标准实践--nz100100噪声向量维度改小会降低生成多样性--ngf6464G中base channel数增大提升容量但易过拟合典型启动命令显存8GB场景python train-dcgan.py --epochs 100 --batch_size 32 --lr 0.0001 --save_dir ./checkpoints/执行后会在./checkpoints/生成G_epoch_50.pth/D_epoch_50.pth每50轮保存loss_log.txt三列epoch, D_loss, G_losssamples/目录每10轮生成fake_epoch_XX.png64张图拼成8×8网格3.3 Loss监控与收敛判断拒绝“看图玄学”用数据说话不要只盯着samples/fake_epoch_XX.png是否“像图”先看loss_log.txt健康信号D_loss在前20轮快速下降至1.5~2.5G_loss同步缓慢上升至1.0~1.850轮后两者在±0.3内小幅震荡危险信号D_loss 0.3且持续下降 → D过强G无法学习需调小D lr或增D训练步数翻车信号G_loss突然飙升至5.0 → G梯度爆炸检查model.py中nn.LeakyReLU(0.2)是否误写为nn.ReLU()我一般用pandas快速绘图import pandas as pd import matplotlib.pyplot as plt df pd.read_csv(./checkpoints/loss_log.txt, sep,, names[epoch,D_loss,G_loss]) plt.plot(df[epoch], df[D_loss], labelD Loss) plt.plot(df[epoch], df[G_loss], labelG Loss) plt.legend(); plt.xlabel(Epoch); plt.ylabel(Loss); plt.grid() plt.savefig(./checkpoints/loss_curve.png)注意loss_log.txt是追加写入若中断重训需手动清空该文件否则曲线错乱。4. 避坑指南98分项目背后的5个血泪经验省下你3天debug时间4.1 现象train-dcgan.py报错ModuleNotFoundError: No module named torchvision原因项目依赖torchvision.transforms做图像增强但pip install torch不自动安装torchvision且版本必须严格匹配PyTorch 1.8.1 → torchvision 0.9.1解决执行pip install torchvision0.9.1cu111 -f https://download.pytorch.org/whl/torch_stable.html务必带cu111后缀否则CPU版torchvision会与CUDA版PyTorch冲突。4.2 现象训练初期D_loss0.000G_lossinf生成图全黑原因model.py中判别器最后一层nn.Linear(512,1)输出未经过nn.Sigmoid()而BCELoss要求输入在[0,1]区间。原项目用nn.BCEWithLogitsLoss()自动sigmoidlog但若误换成nn.BCELoss()就会崩溃。解决检查train-dcgan.py第112行损失函数定义确认是criterion nn.BCEWithLogitsLoss()。若需改用BCELoss则在D输出后加scores torch.sigmoid(scores)。4.3 现象simple-distributions.py生成图是纯色块无纹理原因生成器G的权重未正确初始化。model.py第42行nn.init.normal_(m.weight.data, 0.0, 0.02)被注释或删除导致ConvTranspose2d权重全零。解决打开model.py找到def weights_init(m):函数确保if isinstance(m, nn.ConvTranspose2d):分支内的nn.init.normal_未被注释且0.02标准差未被改为0。4.4 现象utils.py加载图像时报cv2.error: OpenCV(4.5.5) ... invalid value in function cv::resize原因输入图像存在损坏如EXIF旋转标记未处理或尺寸小于64×64cv2.resize无法缩放。解决在load_data()函数中cv2.resize前加校验if img.shape[0] 64 or img.shape[1] 64: img cv2.resize(img, (128,128), interpolationcv2.INTER_CUBIC) # 先放大再裁 img cv2.resize(img, (64,64))4.5 现象训练到50轮后loss突变生成图出现大量条纹伪影原因ops.py中gradient_penalty计算时eps插值系数固定为0.1但在高分辨率或大batch下该值导致梯度惩罚过强。解决修改ops.py第68行eps torch.rand(real_img.size(0), 1, 1, 1)为eps torch.rand(real_img.size(0), 1, 1, 1, devicedevice)补上device参数否则混合精度训练时eps在CPU而real_img在GPU触发隐式拷贝错误。5. 模型部署与效果增强三步让修复结果从“能看”到“可用”5.1 修复单张图像脱离训练框架的轻量推理脚本毕设答辩常被问“能修我这张图吗”此时需一个独立推理脚本。新建inference.pyimport torch import cv2 import numpy as np from model import Generator # 1. 加载训练好的生成器 G Generator(nz100, ngf64, nc3) G.load_state_dict(torch.load(./checkpoints/G_epoch_100.pth)) G.eval() # 关闭dropout/batchnorm # 2. 读取待修复图并预处理 img cv2.imread(./input/test.jpg)[:, :, ::-1] # BGR→RGB img cv2.resize(img, (64,64)) / 255.0 * 2 - 1 # 归一化到[-1,1] img torch.from_numpy(img.transpose(2,0,1)).float().unsqueeze(0) # (1,3,64,64) # 3. 生成修复图 with torch.no_grad(): noise torch.randn(1, 100, 1, 1) fake G(noise).squeeze(0).cpu().numpy() # (3,64,64) fake ((fake 1) / 2 * 255).astype(np.uint8).transpose(1,2,0) # [-1,1]→[0,255] cv2.imwrite(./output/repaired.jpg, fake[:, :, ::-1]) # RGB→BGR保存关键点G.eval()必不可少否则BatchNorm统计量会污染unsqueeze(0)补batch维度squeeze(0)移除batch维度以便后续处理。5.2 效果增强后处理三板斧提升视觉可信度GAN生成图常有颜色偏移、边缘锯齿、局部模糊用OpenCV做低成本增强增强类型OpenCV代码作用参数建议白平衡cv2.cvtColor(fake, cv2.COLOR_RGB2LAB)→lab[:,:,0] cv2.equalizeHist(lab[:,:,0])→cv2.cvtColor(lab, cv2.COLOR_LAB2RGB)校正整体色调必做尤其修复老照片边缘锐化kernel np.array([[0,-1,0],[-1,5,-1],[0,-1,0]])→sharpened cv2.filter2D(fake, -1, kernel)强化破损边缘结构kernel可微调5→6增强强度去马赛克fake cv2.resize(fake, (256,256), interpolationcv2.INTER_CUBIC)→fake cv2.resize(fake, (64,64), interpolationcv2.INTER_AREA)抑制高频噪声仅当生成图有明显块效应时启用将三者串联修复图质感提升显著答辩时老师会直观感受到“这不像GAN乱画的”。5.3 毕设加分项可视化中间特征证明你真懂GAN在学什么评审最爱看“为什么有效”。在model.py的Generator中插入hook# 在Generator.__init__末尾添加 self.feature_maps {} def hook_fn(module, input, output): self.feature_maps[module._get_name()] output.detach().cpu().numpy() self.main[0].register_forward_hook(hook_fn) # 第一层ConvTranspose2d然后在inference.py生成fake后打印G.feature_maps[ConvTranspose2d].shape应为(1,1024,4,4)再用matplotlib显示前4个通道import matplotlib.pyplot as plt feats G.feature_maps[ConvTranspose2d][0] # (1024,4,4) fig, axes plt.subplots(2,2) for i, ax in enumerate(axes.flat): ax.imshow(feats[i], cmapviridis) plt.savefig(./output/features.png)这张图能说明G第一层已学会提取低频结构如轮廓、大块色块而非随机噪声——这就是你答辩时说“生成器在早期层捕获全局结构”的证据。从那以后我每次交毕设代码都强制走一遍simple-distributions.py → train-dcgan.py10轮→ inference.py → features.py四步验证链确保从噪声生成、训练收敛、单图推理到原理可视化全部闭环。这比单纯调参重要十倍——因为答辩时老师问的从来不是“loss多少”而是“你观察到了什么现象怎么解释它”。希望帮到你。本文还有配套的精品资源点击获取
返回列表