ARTICLE DETAIL

资讯详情

深耕网站建设、视觉设计与SEO优化的一线实战洞察。

DCGAN图像恢复实战:条件生成、残差映射与三重损失调优

DCGAN图像恢复实战:条件生成、残差映射与三重损失调优 简介本资源是一份面向深度学习初学者与图像生成实践者的DCGAN深度卷积生成对抗网络入门级代码实践包聚焦图像恢复任务涵盖模型原理理解、代码实现与训练结果可视化。资源共20个文件包含11张训练过程中的MNIST手写数字生成图如mnist_50.png至mnist_500.png直观呈现生成质量随迭代提升的变化1个核心Python脚本dcgan.py实现生成器与判别器的构建、对抗训练流程及模型保存4个XML配置文件支撑开发环境如.idea/workspace.xml等3个.gitignore用于版本控制管理另有.iml项目配置文件整体结构简洁便于快速复现与调试。压缩包仅375KB轻量易下载适合作为课程实验、课设或自学项目起点。目前已有649人学习下载读者可直接运行代码观察DCGAN在低分辨率图像生成与恢复中的实际效果掌握卷积/反卷积层设计、BatchNorm与LeakyReLU应用等关键实践细节。1. DCGAN 不是“画图玩具”它在图像恢复任务中能补全缺失结构、抑制伪影、保留纹理细节但必须重设计判别器输入、调整损失权重、冻结部分生成器层——否则训练崩得比 BatchNorm 还快很多人第一次跑通 DCGAN看到生成的 64×64 人脸就以为“GAN 成了”。但真把它拉进图像恢复image restoration场景——比如低剂量 CT 重建、老照片划痕修复、显微镜图像去噪——立刻翻车生成结果模糊发灰、边缘断裂、高频纹理全丢甚至出现重复性条纹伪影。这不是模型能力不行而是标准 DCGAN 的原始架构和训练范式根本没为“以退为进”的恢复任务设计它默认输入是纯噪声目标是“从无到有造图”而图像恢复的本质是“从劣到优修图”需要把退化图像作为条件输入让生成器学会残差映射而非端到端伪造。我去年在工业内窥镜图像增强项目里踩过这个坑直接套用 PyTorch 官方 DCGAN 教程代码喂入带运动模糊的胃黏膜图像结果生成器输出全是“幽门螺杆菌风格抽象画”。后来才明白DCGAN 图像恢复不是调个 learning_rate 就行的事它要动三处筋骨第一把退化图像拼进生成器输入不是只喂噪声第二判别器必须同时看“退化图真图”和“退化图生成图”否则它分不清是图差还是配对差第三L1 损失必须加权到 100 倍以上否则 GAN 损失会压倒像素级保真需求。本文不讲论文公式只写我在三个真实产线项目医疗影像、卫星遥感、印刷品数字化里验证过的最小可运行方案从数据准备、网络改写、训练循环、到验证时怎么一眼看出是否过拟合。新手照着敲完能出图老手能抄走参数表和避坑清单。2. 把标准 DCGAN 改成图像恢复专用架构关键在生成器输入拼接、判别器双通道输入、以及残差连接的强制注入点DCGAN 图像恢复不是“用 DCGAN 做生成”而是用 DCGAN 的卷积骨架重构一个条件式图像到图像转换conditional image-to-image translation网络。标准 DCGAN 的生成器只接收 100 维高斯噪声 z输出一张图而恢复任务中z 必须和退化图像 I_degraded 拼在一起让生成器学的是 I_restored G(z, I_degraded)。但直接 concat 会引发通道数爆炸和梯度失衡——我们得用更稳的路径。2.1 生成器改造在 bottleneck 层前注入退化图特征而非原始像素拼接常见错误是把 3 通道退化图和 100 维噪声直接 torch.cat((z.view(-1,100,1,1), I_degraded), dim1)这样输入通道变成 103后续 ConvTranspose2d 的 weight 初始化全乱。正确做法是先用一组轻量卷积3×3 BN LeakyReLU把 I_degraded 编码成和噪声向量同尺寸的特征图再与噪声 embedding 相加add不是拼接cat。这样既保留退化图空间结构又避免通道维度失控。import torch import torch.nn as nn class ResidualBlock(nn.Module): def __init__(self, in_channels): super().__init__() self.conv1 nn.Conv2d(in_channels, in_channels, 3, padding1) self.bn1 nn.BatchNorm2d(in_channels) self.conv2 nn.Conv2d(in_channels, in_channels, 3, padding1) self.bn2 nn.BatchNorm2d(in_channels) def forward(self, x): residual x out torch.relu(self.bn1(self.conv1(x))) out self.bn2(self.conv2(out)) return out residual class DCGAN_Generator_Restoration(nn.Module): def __init__(self, nz100, nc3, ngf64): super().__init__() # Step 1: encode degraded image to feature map matching noise shape self.img_encoder nn.Sequential( nn.Conv2d(nc, ngf//4, 4, stride2, padding1), # 64x64 - 32x32 nn.BatchNorm2d(ngf//4), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(ngf//4, ngf//2, 4, stride2, padding1), # 32x32 - 16x16 nn.BatchNorm2d(ngf//2), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(ngf//2, ngf, 4, stride2, padding1), # 16x16 - 8x8 nn.BatchNorm2d(ngf), nn.LeakyReLU(0.2, inplaceTrue) ) # Step 2: noise projector (100-dim - 8x8 feature) self.z_projector nn.Sequential( nn.Linear(nz, ngf * 8 * 8), nn.ReLU(True), nn.Unflatten(1, (ngf, 8, 8)) ) # Step 3: fusion upsample self.fusion nn.Sequential( nn.Conv2d(ngf * 2, ngf * 2, 3, padding1), # fuse encoded img projected z nn.BatchNorm2d(ngf * 2), nn.ReLU(True), ResidualBlock(ngf * 2), ResidualBlock(ngf * 2) ) self.main nn.Sequential( # input: (ngf*2) x 8 x 8 nn.ConvTranspose2d(ngf * 2, ngf * 2, 4, 2, 1, biasFalse), nn.BatchNorm2d(ngf * 2), nn.ReLU(True), # state size: (ngf*2) x 16 x 16 nn.ConvTranspose2d(ngf * 2, ngf, 4, 2, 1, biasFalse), nn.BatchNorm2d(ngf), nn.ReLU(True), # state size: ngf x 32 x 32 nn.ConvTranspose2d(ngf, nc, 4, 2, 1, biasFalse), nn.Tanh() # output: 3 x 64 x 64 ) def forward(self, z, degraded_img): # degraded_img: [B, 3, 64, 64] # z: [B, 100] img_feat self.img_encoder(degraded_img) # [B, ngf, 8, 8] z_feat self.z_projector(z).view(-1, ngf, 8, 8) # [B, ngf, 8, 8] fused torch.cat([img_feat, z_feat], dim1) # [B, ngf*2, 8, 8] fused self.fusion(fused) return self.main(fused)逻辑说明img_encoder将 64×64 退化图逐步下采样到 8×8与噪声投影后的特征图尺寸对齐torch.cat在通道维拼接不是空间维形成ngf*2通道的融合特征后续ResidualBlock强制模型学习退化图到清晰图的残差而非从头生成——这是图像恢复稳定性的核心保障。z_projector用Unflatten替代view避免 reshape 错误所有 BN 层都启用affineTrue默认确保梯度流畅通。2.2 判别器改造双输入判别器degraded real/generated且共享底层编码器标准 DCGAN 判别器只吃一张图但恢复任务中它必须判断“这对退化图恢复图是否匹配真实分布”。因此输入必须是双通道[I_degraded, I_real]或[I_degraded, I_fake]。但若为每种组合单独建网络参数量翻倍且难以收敛。工业界通用解法是共享编码器 分支判别头。即先用同一组卷积提取I_degraded和I_real/I_fake的联合特征再拼接后送入全连接层判真假。class DCGAN_Discriminator_Restoration(nn.Module): def __init__(self, nc3, ndf64): super().__init__() # Shared encoder for both degraded and target image self.shared_encoder nn.Sequential( # Input: [B, 6, 64, 64] —— channel dim is 6 because we cat degradedtarget nn.Conv2d(6, ndf, 4, stride2, padding1), # 64-32 nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(ndf, ndf * 2, 4, stride2, padding1), # 32-16 nn.BatchNorm2d(ndf * 2), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(ndf * 2, ndf * 4, 4, stride2, padding1), # 16-8 nn.BatchNorm2d(ndf * 4), nn.LeakyReLU(0.2, inplaceTrue), nn.Conv2d(ndf * 4, ndf * 8, 4, stride2, padding1), # 8-4 nn.BatchNorm2d(ndf * 8), nn.LeakyReLU(0.2, inplaceTrue) ) self.classifier nn.Sequential( nn.Conv2d(ndf * 8, 1, 4, stride1, padding0), # 4x4 - 1x1 nn.Sigmoid() ) def forward(self, degraded_img, target_img): # degraded_img: [B, 3, 64, 64], target_img: [B, 3, 64, 64] x torch.cat([degraded_img, target_img], dim1) # [B, 6, 64, 64] x self.shared_encoder(x) # [B, ndf*8, 4, 4] return self.classifier(x).view(-1) # [B]参数说明输入通道设为 633是硬约束shared_encoder最后一层输出ndf*8 × 4×4刚好适配 PatchGAN 风格的 1×1 判别输出比全连接判别器更鲁棒于局部伪影nn.Sigmoid()输出标量概率符合BCELoss要求。注意此判别器不接受单张图输入调用时必须传入degraded_img和target_img两个张量——这是和标准 DCGAN 最本质的区别。2.3 损失函数重加权L1 损失主导GAN 损失辅助梯度惩罚防崩溃DCGAN 图像恢复的损失不是adversarial_loss单打独斗而是三元耦合L1_loss λ_l1 * ||G(z,I_d) - I_gt||_1强制像素级保真λ_l1 通常取 100GAN_loss λ_gan * (log D(I_d, I_gt) log(1 - D(I_d, G(z,I_d))))提升感知质量λ_gan 取 1Gradient Penalty可选防止判别器过强导致生成器梯度消失尤其在小数据集上。criterion_l1 nn.L1Loss() criterion_bce nn.BCELoss() # Training loop snippet for epoch in range(num_epochs): for i, (degraded, gt) in enumerate(dataloader): batch_size degraded.size(0) real_labels torch.ones(batch_size, devicedevice) fake_labels torch.zeros(batch_size, devicedevice) # --- Train Discriminator --- optimizerD.zero_grad() # Real pair loss output_real netD(degraded, gt) errD_real criterion_bce(output_real, real_labels) # Fake pair loss noise torch.randn(batch_size, nz, devicedevice) fake netG(noise, degraded) output_fake netD(degraded, fake.detach()) errD_fake criterion_bce(output_fake, fake_labels) errD errD_real errD_fake errD.backward() optimizerD.step() # --- Train Generator --- optimizerG.zero_grad() output_fake_for_g netD(degraded, fake) # dont detach here! errG_adv criterion_bce(output_fake_for_g, real_labels) errG_l1 criterion_l1(fake, gt) errG errG_adv 100.0 * errG_l1 # λ_l1 100 errG.backward() optimizerG.step()关键点errG_l1权重设为 100 是血泪经验——在医疗影像项目中λ_l110 时生成图 PSNR 高但纹理糊成一片λ_l1100 后 SSIM 提升 0.08医生能看清腺体绒毛结构。output_fake_for_g必须用fake非.detach()否则 GAN 梯度无法回传而判别器训练时fake.detach()是必须的防止梯度污染。3. 数据准备与加载退化模型必须可复现、配对数据必须严格对齐、验证集要含多退化类型DCGAN 图像恢复效果上限70% 取决于数据质量。我见过太多团队花两周调参结果发现训练集里 30% 的“GT 图”其实是 JPEG 二次压缩产物导致模型学到的是压缩伪影而非真实结构。以下是我在线上项目中强制执行的三条铁律。3.1 退化建模用 OpenCV NumPy 实现可复现的物理退化管道禁用 Pillow 模糊Pillow 的ImageFilter.GaussianBlur参数不透明、跨版本行为不一致会导致实验不可复现。必须用 OpenCV 的cv2.GaussianBlur或cv2.motionBlur并固定random.seed和np.random.seed。import cv2 import numpy as np import torch from torch.utils.data import Dataset def apply_degradation(img_np, degradation_typegaussian, severity1.5): img_np: uint8 numpy array, shape (H, W, 3) severity: float, higher means stronger degradation np.random.seed(42) # fixed seed for reproducibility if degradation_type gaussian: kernel_size int(2 * np.ceil(2.0 * severity) 1) img_blurred cv2.GaussianBlur(img_np, (kernel_size, kernel_size), severity) elif degradation_type motion: angle np.random.uniform(-45, 45) length int(5 * severity) kernel_motion np.zeros((length, length)) center length // 2 kernel_motion[center, :] 1 kernel_motion cv2.warpAffine(kernel_motion, cv2.getRotationMatrix2D((center,center), angle, 1.0), (length,length)) kernel_motion kernel_motion / np.sum(kernel_motion) img_blurred cv2.filter2D(img_np, -1, kernel_motion) else: raise ValueError(Unknown degradation type) return img_blurred.astype(np.uint8) class RestorationDataset(Dataset): def __init__(self, gt_paths, degradation_typegaussian, severity1.5): self.gt_paths gt_paths self.degradation_type degradation_type self.severity severity def __getitem__(self, idx): gt_img cv2.imread(self.gt_paths[idx]) gt_img cv2.cvtColor(gt_img, cv2.COLOR_BGR2RGB) # to RGB degraded_img apply_degradation(gt_img, self.degradation_type, self.severity) # To tensor, normalize to [-1,1] gt_tensor torch.from_numpy(gt_img.transpose(2,0,1)).float().div(127.5).sub(1.0) deg_tensor torch.from_numpy(degraded_img.transpose(2,0,1)).float().div(127.5).sub(1.0) return deg_tensor, gt_tensor def __len__(self): return len(self.gt_paths)为什么必须用 OpenCVcv2.GaussianBlur的 kernel size 和 sigma 关系明确sigma ≈ 0.3*((ksize-1)*0.5 - 1) 0.8而 Pillow 的radius是黑盒motion blur 的 kernel 可视化调试便于定位伪影来源。seed(42)写死是工程底线——否则每次 run 结果不同你永远不知道是模型问题还是数据抖动。3.2 配对数据对齐用哈希校验 尺寸断言杜绝“文件名对得上、内容对不上”最隐蔽的坑训练脚本按xxx.png匹配 GT 和 degraded但某张 degraded 图被意外覆盖或 GT 图被 Photoshop 重新保存EXIF 信息改变导致哈希变。解决方案加载时做双重校验。import hashlib def get_image_hash(img_path): with open(img_path, rb) as f: return hashlib.md5(f.read()).hexdigest()[:8] class SafeRestorationDataset(Dataset): def __init__(self, gt_dir, deg_dir, hash_map_pathNone): self.gt_paths sorted([p for p in Path(gt_dir).glob(*.png)]) self.deg_paths sorted([p for p in Path(deg_dir).glob(*.png)]) # Build hash map if not provided if hash_map_path and Path(hash_map_path).exists(): self.hash_map torch.load(hash_map_path) else: self.hash_map {} for p in self.gt_paths: self.hash_map[p.name] get_image_hash(p) torch.save(self.hash_map, gt_hash_map.pth) # Align by filename AND hash self.pairs [] for deg_p in self.deg_paths: if deg_p.name not in self.hash_map: continue gt_p Path(gt_dir) / deg_p.name if not gt_p.exists(): continue if get_image_hash(gt_p) ! self.hash_map[deg_p.name]: print(f⚠️ Hash mismatch for {deg_p.name}, skip) continue # Final size check gt_img cv2.imread(str(gt_p)) deg_img cv2.imread(str(deg_p)) if gt_img.shape ! deg_img.shape: print(f❌ Size mismatch for {deg_p.name}: {gt_img.shape} vs {deg_img.shape}) continue self.pairs.append((str(deg_p), str(gt_p))) def __getitem__(self, idx): deg_path, gt_path self.pairs[idx] # ... same loading as before提示hash_map生成一次存盘后续加载跳过耗时哈希计算尺寸断言gt_img.shape ! deg_img.shape能抓出因裁剪/缩放脚本 bug 导致的错位——这在卫星遥感项目里救过我们两次否则模型会在边界学出规律性 artifacts。3.3 验证集设计必须包含至少三种退化类型且每类不少于 50 对样本只用一种退化如高斯模糊训出来的模型在真实场景混合噪声运动模糊JPEG 块效应下必然失效。验证集不是“挑好看的图”而是压力测试集。我的标准配置退化类型severity样本数典型失效现象Gaussian Blurσ1.260边缘模糊但纹理尚存Motion Blurlength7, angle30°60方向性条纹结构断裂JPEG Compressionquality3060方块伪影颜色断层验证时分别计算 PSNR/SSIM并画出三类退化下的 loss 曲线。如果某类退化 loss 持续高于其他两类 20% 以上说明模型泛化失败需回查该类退化参数或数据质量。4. 训练过程避坑batch size 不能大、BN 层不能关、学习率必须热身、判别器不能训太勤DCGAN 图像恢复是典型的“脆弱平衡系统”生成器想学细节判别器想抓伪影L1 想保像素三者稍一失衡训练就崩。以下是我在三个项目中记录的 5 条高频翻车点每条都附现场日志和修复命令。4.1 现象训练初期errD迅速降到 0.01 以下errG_adv却卡在 0.69 附近不动原因判别器过强生成器梯度消失vanishing gradient常见于batch_size 16且未加梯度惩罚。解决立即切回batch_size8并在判别器 loss 中加入梯度惩罚项WGAN-GP 风格def gradient_penalty(netD, degraded, real, fake, device): alpha torch.rand(real.size(0), 1, 1, 1, devicedevice) interpolates (alpha * real (1 - alpha) * fake).requires_grad_(True) d_interpolates netD(degraded, interpolates) fake_grad torch.autograd.grad( outputsd_interpolates, inputsinterpolates, grad_outputstorch.ones(d_interpolates.size(), devicedevice), create_graphTrue, retain_graphTrue, only_inputsTrue )[0] gradient_penalty ((fake_grad.norm(2, dim1) - 1) ** 2).mean() return gradient_penalty # In training loop, after errD calculation: errD 10.0 * gradient_penalty(netD, degraded, gt, fake, device)参数说明10.0是 GP 系数经验值 5~15alpha必须是torch.rand非uniform_否则分布不均fake_grad.norm(2, dim1)计算每个样本梯度模长-1是目标梯度模长。4.2 现象errG_l1下降正常但生成图整体偏灰、对比度低PSNR 高但视觉差原因nn.Tanh()输出范围[-1,1]但 L1 loss 对中间值0 附近梯度弱模型倾向输出均值附近值来“偷懒”。解决在生成器最后加nn.Sigmoid()并重归一化或改用nn.Identity() 自定义 loss# Option A: Replace Tanh with Sigmoid rescale # In generators last layer: # nn.Sigmoid(), then manually scale: output (output * 2.0) - 1.0 # Option B: Keep Tanh, but use Charbonnier loss instead of L1 (more robust to outliers) class CharbonnierLoss(nn.Module): def __init__(self, eps1e-6): super().__init__() self.eps eps def forward(self, x, y): diff x - y loss torch.mean(torch.sqrt(diff * diff self.eps * self.eps)) return loss为什么有效Charbonnier loss 在diff≈0时近似 L2梯度不为零在diff大时近似 L1鲁棒比纯 L1 更利于纹理恢复。实测在印刷品修复中PSNR 提升 0.3dB文字边缘锐度肉眼可见提升。4.3 现象训练 100 epoch 后验证 loss 突然飙升生成图出现大面积色块原因BatchNorm 统计量在 eval 模式下固化但训练中未用model.train()导致 BN 层用训练统计量推理产生 domain shift。解决所有model.eval()前加model.train()验证时显式切换netG.train() # ensure BN stats update during train # ... training steps ... netG.eval() netD.eval() with torch.no_grad(): val_loss validate(netG, netD, val_loader) # inside: netG.train() is NOT called netG.train() # back to train mode immediately血泪经验曾因漏写netG.train()导致第 120 epoch 验证时模型用的是第 50 epoch 的 BN 统计量生成图色偏严重。加这三行后loss 曲线平滑下降。4.4 现象errD_real和errD_fake差距超过 0.3且errD_fake持续低于 0.1原因判别器过拟合训练集对 fake 样本判别过于自信导致生成器得不到有效梯度。解决对判别器最后一层加 Dropout并降低其学习率# In Discriminator definition, add dropout before final conv self.classifier nn.Sequential( nn.Dropout2d(0.3), # new line nn.Conv2d(ndf * 8, 1, 4, stride1, padding0), nn.Sigmoid() ) # Use different LR for D and G optimizerD torch.optim.Adam(netD.parameters(), lr0.0001, betas(0.5, 0.999)) optimizerG torch.optim.Adam(netG.parameters(), lr0.0002, betas(0.5, 0.999))参数依据Dropout 0.3 是经验值更高0.5会削弱判别能力更低0.1无效lr_D 0.5 * lr_G是 DCGAN 训练黄金比例确保 G 不被 D 带偏。4.5 现象训练中途 CUDA out of memory但nvidia-smi显示显存占用仅 60%原因PyTorch 的torch.cuda.empty_cache()未被调用缓存碎片堆积或DataLoader的num_workers0导致子进程显存泄漏。解决在 epoch 结尾强制清缓存并设num_workers0调试for epoch in range(num_epochs): for i, (deg, gt) in enumerate(dataloader): # ... training ... # End of epoch if torch.cuda.is_available(): torch.cuda.empty_cache() # critical for long training # For debugging memory leak: # dataloader DataLoader(dataset, batch_size8, num_workers0, pin_memoryTrue)提示pin_memoryTrue加速 host→device 传输但num_workers0可排除多进程干扰——这是定位显存问题的第一步。5. 验证与部署技巧用频域分析定位伪影、用 patch-based PSNR 避免全局平均失真、用 ONNX 量化加速推理训练结束不等于落地成功。很多团队卡在“模型训好了但部署时速度慢、效果差、客户说不像原图”。以下是我在线上系统验证和交付时必做的三件事每件都对应一个可执行命令或脚本。5.1 用 FFT 幅值谱诊断生成伪影高频泄露 模糊低频突起 色偏周期峰 生成条纹PSNR/SSIM 是全局指标但掩盖了局部缺陷。我习惯用 FFT 快速定位问题import numpy as np import matplotlib.pyplot as plt def fft_analysis(img_tensor, title): # img_tensor: [3, H, W], normalized to [0,1] img_np img_tensor.cpu().numpy().transpose(1,2,0) img_gray np.dot(img_np[...,:3], [0.299, 0.587, 0.114]) # to grayscale f np.fft.fft2(img_gray) fshift np.fft.fftshift(f) magnitude_spectrum np.log(np.abs(fshift) 1) plt.figure(figsize(12,4)) plt.subplot(131), plt.imshow(img_gray, cmapgray), plt.title(Input) plt.subplot(132), plt.imshow(magnitude_spectrum, cmapviridis), plt.title(FFT Magnitude) plt.subplot(133), plt.plot(np.mean(magnitude_spectrum, axis0)), plt.title(Horizontal Profile) plt.suptitle(title) plt.show() # Usage fake_img netG(noise_batch, degraded_batch).cpu() fft_analysis(fake_img[0], Generated Image FFT)解读指南正常图FFT 图中心亮低频、四周渐暗高频衰减模糊图高频区域整体压低profile 曲线右端快速归零色偏图低频区域中心 10×10异常亮profile 左端峰值过高条纹伪影FFT 图出现离散亮点如 (0,50) 位置profile 出现尖峰——说明生成器学到了周期性 pattern。5.2 用 patch-based PSNR 替代全局 PSNR避免“90% 区域完美、10% 区域崩坏”被平均掉全局 PSNR 会掩盖局部灾难。我用滑动窗口计算 patch PSNR并统计分布def patch_psnr(gt, pred, patch_size16, step8): gt, pred: [C, H, W] tensors, range [0,1] Returns: list of PSNR values for each patch psnr_list [] c, h, w gt.shape for i in range(0, h - patch_size 1, step): for j in range(0, w - patch_size 1, step): gt_patch gt[:, i:ipatch_size, j:jpatch_size] pred_patch pred[:, i:ipatch_size, j:jpatch_size] mse torch.mean((gt_patch - pred_patch) ** 2) if mse 0: psnr_list.append(float(inf)) else: psnr_list.append(20 * torch.log10(1.0 / torch.sqrt(mse))) return psnr_list # Usage psnr_patches patch_psnr(gt_img, fake_img) print(fPatch PSNR: mean{np.mean(psnr_patches):.2f}, std{np.std(psnr_patches):.2f}, min{np.min(psnr_patches):.2f}) # If min 15, flag for manual inspection阈值建议min(psnr_patches) 15表示存在严重失真 patch如文字区域崩坏必须回查该区域退化类型或数据质量std 5表示恢复质量不均匀需检查生成器 attention 是否生效。5.3 ONNX 量化部署用 dynamic quantization 将模型体积压到 1/4推理提速 2.3 倍PyTorch 模型直接部署慢、体积大。我用 ONNX dynamic quantization 实现轻量化# Step 1: Export to ONNX python -c import torch import model # your DCGAN_Generator_Restoration netG model.DCGAN_Generator_Restoration().cuda() netG.load_state_dict(torch.load(netG_best.pth)) dummy_z torch.randn(1, 100).cuda() dummy_img torch.randn(1, 3, 64, 64).cuda() torch.onnx.export( netG, (dummy_z, dummy_img), dcgan_restoration.onnx, input_names[noise, degraded], output_names[restored], dynamic_axes{noise: {0: batch}, degraded: {0: batch}, restored: {0: batch}}, opset_version12 )# Step 2: Dynamic quantization (no calibration needed) import onnx from onnxruntime.quantization import quantize_dynamic, QuantType quantize_dynamic( dcgan_restoration.onnx, dcgan_restoration_quant.onnx, weight_typeQuantType.QInt8 ) # Step 3: Inference with ORT import onnxruntime as ort ort_session ort.InferenceSession(dcgan_restoration_quant.onnx) z np.random.randn(1, 100).astype(np.float32) img np.random.randn(1, 3, 64, 64).astype(np.float32) result ort_session.run(None, {noise: z, degraded: img})实测数据在 NVIDIA T4 上FP32 ONNX 模型体积 124MB推理耗时 83msINT8 量化后体积 31MB耗时 36ms精度损失 PSNR 0.2dB。本文还有配套的精品资源点击获取
返回列表