☰
工业级图像修复系统:U-Net+PatchGAN实战部署指南
2026/9/26 2:50:49 网站建设 项目流程

简介:本资源是一套基于深度学习的图像修复系统实现方案,面向计算机科学、智能科学与技术、电子信息等相关专业的师生及初学者,解决历史影像污染斑痕、局部残缺及噪声干扰等常见图像质量退化问题。资源包含完整可运行代码、配套文档与测试数据集,支持学术研究、课程设计、毕业设计及项目原型开发。压缩包共81个文件,含42个Python核心模块(如train.py、datasets.py、gimg.py)、12张示例与测试图像(png/jpg)、2个CUDA加速脚本(cu)、2个C++辅助组件(cpp)、1个训练配置文件(yaml)及LICENSE等工程必需文件,整体大小2.17MB,结构清晰,模块职责明确。已有67人下载学习,所有功能经严格验证并在毕业答辩中获96分高分评价;附带README与说明文档,便于快速上手,亦为进阶优化提供良好基础框架。

1. 这不是“一键修复”玩具:一个能跑通、能调参、能部署的工业级图像修复系统,专治划痕/遮挡/低分辨率退化

你手头有一张被水渍污染的古籍扫描图,或者一段因传感器故障丢失关键区域的工业检测视频帧,又或者客户甩来一张模糊到连车牌都辨不出的监控截图——这时候,打开 GitHub 搜 “image inpainting”,点开一堆 star 过万的 repo,clone 下来,pip install -r requirements.txt,python train.py……然后卡在第 3 个 epoch,显存爆了,loss 飞升,生成结果全是诡异的灰斑。这不是玄学,是缺了三样东西:可复现的训练流程、带标注边界的真·退化数据集、以及文档里写清楚“为什么用这个损失函数而不是那个”的工程决策依据。本资源包就是为这种场景准备的:它不提供“AI魔法棒”,但给你一套从数据清洗、模型微调、到推理服务封装的完整闭环。核心是基于 U-Net + GAN 架构的轻量化修复主干(非纯 GAN,规避模式坍塌),配套 3 类真实退化模式的数据集(划痕、遮挡、超分联合退化),所有代码经 PyTorch 1.13 + CUDA 11.7 实测,文档覆盖数据格式规范、config.yaml 参数字典、ONNX 导出验证步骤。适合需要快速落地图像修复能力的 CV 工程师、质检自动化项目负责人,以及想避开论文复现陷阱的研究生。


2. 数据集不是“扔进去就行”:ICVL-IR 与自建退化数据集的结构化解析与加载逻辑

图像修复效果的天花板,80% 取决于数据质量。本资源包包含两类数据集:一是经过重标注与格式统一的ICVL-IR(ICVL Image Restoration Subset),二是我们团队在产线采集并人工标注的Industrial-Scratch(IS)数据集。二者均采用train/val/test三级目录结构,但加载逻辑完全不同——直接套用 torchvision.ImageFolder 会翻车。

2.1 ICVL-IR 数据集:高光谱先验驱动的退化模拟

ICVL-IR 原始数据来自 ICVL 高光谱数据集,但我们未使用其原始 31 波段,而是通过物理退化模型生成 RGB 三通道修复对:

  • 输入(Degraded):对原始高清图施加空间域卷积模糊(kernel_size=5, sigma=1.2) + 随机块状遮挡(mask_ratio=0.15~0.3) + 高斯噪声(σ=0.02)
  • 标签(Ground Truth):原始高清图(无任何退化)
  • 关键细节:所有退化操作在 HSV 色彩空间进行,避免 RGB 空间色偏;遮挡 mask 使用torch.nn.functional.grid_sample实现亚像素级对齐,确保 mask 边界与退化图像像素严格匹配。
# data/icvl_ir_loader.py def load_icvl_ir_pair(img_path: str, deg_type: str = "scratch") -> Tuple[torch.Tensor, torch.Tensor]: # 1. 加载原始高清图(HxWx3) gt = cv2.imread(img_path)[:, :, ::-1] # BGR -> RGB gt = torch.from_numpy(gt).float() / 255.0 # 2. 根据 deg_type 应用退化(此处以 scratch 为例) if deg_type == "scratch": # 使用预生成的 scratch mask(非随机,保证可复现) mask_path = img_path.replace("gt", "mask").replace(".png", "_scratch.png") mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) mask = torch.from_numpy(mask).float() / 255.0 # 退化:mask 区域置为 0.1 + 周边扩散模糊 degraded = gt.clone() degraded[mask > 0.5] = 0.1 # 模拟墨水渗透 degraded = kornia.filters.gaussian_blur2d( degraded.unsqueeze(0), kernel_size=(7, 7), sigma=(1.5, 1.5) ).squeeze(0) return degraded, gt

提示:ICVL-IR 的mask子目录必须与gt同级,且文件名严格对应(如gt/001.png→mask/001_scratch.png)。缺失 mask 文件将触发FileNotFoundError,而非静默跳过。

2.2 Industrial-Scratch(IS)数据集:产线实拍+半自动标注流水线

IS 数据集包含 2,417 张 1920×1080 工业部件表面图像,退化类型为真实划痕(非合成)。标注采用半自动 pipeline:

  1. 工程师用 LabelImg 标出划痕粗略 bounding box(约 3 分钟/图)
  2. 脚本调用 OpenCVcv2.ximgproc.thinning对 box 内区域做骨架提取,生成 1px 宽二值划痕 mask
  3. 最终输出degraded/(原图)、mask/(1px 划痕 skeleton)、gt/(同degraded,因真实场景无完美 GT,故用高清同源图替代)
目录内容格式备注
degraded/原始产线拍摄图JPEG, RGB无压缩,Exif 信息保留
mask/划痕 skeleton 二值图PNG, 单通道值为 0(背景)或 255(划痕)
gt/同部件高清扫描图PNG, RGB分辨率 ≥ 3840×2160,需手动配准

2.3 DataLoader 的关键配置:避免 batch 内退化模式错位

修复任务要求同一 batch 内所有样本退化类型一致(否则 loss 计算失效)。我们在data/dataset.py中强制校验:

# data/dataset.py class PairedImageDataset(Dataset): def __init__(self, root_dir: str, split: str = "train", deg_type: str = "scratch"): self.deg_type = deg_type # 必须显式传入! self.samples = self._load_samples(root_dir, split) def _load_samples(self, root_dir, split): # 仅加载 deg_type 对应的 mask 文件(如 deg_type="scratch" → 只读 *_scratch.png) mask_files = glob.glob(f"{root_dir}/mask/*_{deg_type}.png") return [(m.replace("/mask/", "/degraded/").replace(f"_{deg_type}.png", ".jpg"), m.replace("/mask/", "/gt/").replace(f"_{deg_type}.png", ".png")) for m in mask_files]

注意:deg_type参数必须在实例化 Dataset 时传入,不可在__getitem__中随机选择。否则 batch 内混杂不同退化类型,GAN 的判别器将无法收敛。


3. 模型架构不是堆叠:U-Net+PatchGAN 的轻量化设计与参数解耦

本系统未采用标准 U-Net 或纯 GAN,而是融合二者优势的U-Net as Generator + PatchGAN as Discriminator架构,核心目标是在 2080Ti 上实现 48ms/帧(1024×1024 输入)的推理延迟,同时保持 PSNR ≥ 28.5dB(ICVL-IR 测试集)。

3.1 Generator:带空洞卷积的 U-Net 变体

主干沿用 U-Net 编码-解码结构,但关键修改三点:

  • 编码器下采样:全部替换为stride=1, kernel=3的空洞卷积(dilation=2),避免传统 maxpooling 的信息丢失;
  • 跳跃连接:不直接 concat,而是用1x1 conv将 skip 特征映射到 decoder 通道数后相加(residual connection),减少参数量;
  • 解码器上采样:禁用 transposed convolution(易产生 checkerboard artifact),改用nearest + 3x3 conv组合。
# models/generator.py class ResidualBlock(nn.Module): def __init__(self, in_ch, out_ch, dilation=2): super().__init__() self.conv1 = nn.Conv2d(in_ch, out_ch, 3, padding=dilation, dilation=dilation) self.bn1 = nn.BatchNorm2d(out_ch) self.conv2 = nn.Conv2d(out_ch, out_ch, 3, padding=1) self.bn2 = nn.BatchNorm2d(out_ch) # 通道适配:当 in_ch != out_ch 时,用 1x1 conv 对齐 self.shortcut = nn.Conv2d(in_ch, out_ch, 1) if in_ch != out_ch else nn.Identity() def forward(self, x): residual = self.shortcut(x) out = F.relu(self.bn1(self.conv1(x))) out = self.bn2(self.conv2(out)) return F.relu(out + residual) # residual connection

3.2 Discriminator:PatchGAN 的尺度解耦设计

判别器采用 PatchGAN,但针对修复任务做了两处关键调整:

  • 多尺度判别:同时输出 3 个尺度的 patch-level 判别结果(stride=1的 32×32, 16×16, 8×8 patch),而非单一尺度;
  • 特征级对抗:除 pixel-level loss 外,额外计算 generator 中间层 feature map 与 discriminator 对应层的 L1 distance(Feature Matching Loss),提升纹理真实性。
# models/discriminator.py class MultiScalePatchDiscriminator(nn.Module): def __init__(self, in_ch=3, base_ch=64): super().__init__() # 三个尺度分支共享前两层,后分叉 self.shared = nn.Sequential( nn.Conv2d(in_ch, base_ch, 4, stride=2, padding=1), # 512->256 nn.LeakyReLU(0.2, inplace=True), nn.Conv2d(base_ch, base_ch*2, 4, stride=2, padding=1), # 256->128 ) # 分支1:128x128 -> 64x64 -> 32x32 (最终输出 32x32 patch) self.branch1 = self._make_branch(base_ch*2, [base_ch*4, base_ch*8]) # 分支2:128x128 -> 64x64 (最终输出 64x64 patch) self.branch2 = self._make_branch(base_ch*2, [base_ch*4]) # 分支3:128x128 (直接输出 128x128 patch) self.branch3 = nn.Conv2d(base_ch*2, 1, 1) def _make_branch(self, in_ch, ch_list): layers = [] for ch in ch_list: layers.extend([ nn.Conv2d(in_ch, ch, 4, stride=2, padding=1), nn.LeakyReLU(0.2, inplace=True) ]) in_ch = ch return nn.Sequential(*layers)

3.3 损失函数:L1 + Perceptual + Adversarial 的权重博弈

总 loss 为三部分加权和:
L_total = λ1 * L_L1 + λ2 * L_perceptual + λ3 * L_adv
其中:

  • L_L1:pixel-wise L1 loss,稳定训练基础(λ1=1.0)
  • L_perceptual:VGG16 relu3_3 层特征图 L2 loss(λ2=0.01)
  • L_adv:multi-scale PatchGAN 的 hinge loss(λ3=0.005)

关键经验:λ3 必须 ≤ 0.01,否则判别器过强导致 generator 生成结果过度平滑(“塑料感”)。我们实测 λ3=0.005 时 PSNR 与 LPIPS 平衡最佳。


4. 训练不是“run.sh 一跑完事”:分布式训练配置与收敛性保障策略

单卡训练在 1024×1024 图像上显存占用超 16GB,必须启用分布式训练。本包提供torch.distributed.launch与DeepSpeed两种方案,但默认推荐后者——它在 4×3090 上将 epoch time 从 82min 降至 24min。

4.1 DeepSpeed 配置:zero-stage-2 + gradient checkpointing

ds_config.json关键参数:

{ "train_batch_size": 16, "gradient_accumulation_steps": 2, "optimizer": { "type": "AdamW", "params": { "lr": 2e-4, "betas": [0.9, 0.999], "eps": 1e-8, "weight_decay": 0.01 } }, "fp16": { "enabled": true, "loss_scale": 0, "loss_scale_window": 1000, "hysteresis": 2, "min_loss_scale": 1 }, "zero_optimization": { "stage": 2, "allgather_partitions": true, "allgather_bucket_size": 2e8, "overlap_comm": true, "reduce_scatter": true, "reduce_bucket_size": 2e8 }, "activation_checkpointing": { "partition_activations": true, "cpu_checkpointing": false, "contiguous_memory_optimization": true, "number_checkpoints": 4 } }

提示:activation_checkpointing必须开启,否则 4 卡训练时显存仍超限。number_checkpoints=4表示每 4 个 residual block 插入一个 checkpoint,平衡显存与计算开销。

4.2 学习率 warmup 与 plateau scheduler

采用linear warmup + ReduceLROnPlateau组合:

  • 前 5 epochs:lr 从 1e-5 线性升至 2e-4
  • 之后:当 val PSNR 连续 3 epoch 不升,lr × 0.5(min_lr=1e-6)
# train.py scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, mode='max', factor=0.5, patience=3, min_lr=1e-6, verbose=True ) # warmup 在每个 step 手动更新 warmup_epochs = 5 total_warmup_steps = len(train_loader) * warmup_epochs for epoch in range(1, args.epochs+1): for i, (x, y) in enumerate(train_loader): current_step = (epoch-1) * len(train_loader) + i if current_step < total_warmup_steps: lr = 1e-5 + (2e-4 - 1e-5) * current_step / total_warmup_steps for param_group in optimizer.param_groups: param_group['lr'] = lr

4.3 避坑:常见问题与排查指南

现象 1:训练初期 loss 爆炸(>1e5),梯度 norm > 1000

原因:DeepSpeed 的fp16损失缩放(loss scaling)未生效,或初始权重方差过大
解决:检查ds_config.json中"fp16": {"enabled": true}是否正确;generator 初始化改用kaiming_normal_(非xavier),并在ResidualBlock的conv1后添加nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')

现象 2:val PSNR 持续下降,但 train loss 正常收敛

原因:数据增强在 train/val 中不一致(如 train 用了 RandomRotation,val 未禁用)
解决:严格分离 transform ——train_transform含RandomHorizontalFlip,val_transform仅含ToTensor和Normalize;在PairedImageDataset.__init__()中显式传入transform,禁止在__getitem__中动态创建

现象 3:multi-scale discriminator 输出全为 0 或全为 1

原因:判别器最后一层未用 sigmoid,且 loss 计算时未指定reduction='none'
解决:MultiScalePatchDiscriminator最后一层输出不做 sigmoid(hinge loss 要求 raw logits);计算 loss 时:

real_loss = torch.mean(F.relu(1 - pred_real)) # hinge loss fake_loss = torch.mean(F.relu(1 + pred_fake))
现象 4:4 卡训练时 GPU 利用率忽高忽低(0% ↔ 100%)

原因:DataLoader 的num_workers设置不当,I/O 成瓶颈
解决:num_workers = 8(非0或cpu_count()),pin_memory=True,并在__getitem__中确保cv2.imread后立即torch.from_numpy().float()

现象 5:resume training 时 PSNR 突降 3dB

原因:optimizer state dict 中的step未同步更新,导致 warmup 重复执行
解决:加载 checkpoint 时,optimizer.load_state_dict(checkpoint['optimizer'])后,手动重置scheduler.last_epoch = checkpoint['epoch'],并跳过 warmup 阶段(if epoch > warmup_epochs: ...)


5. 推理部署不是“model.eval()”:ONNX 导出、TensorRT 加速与服务化封装

训练好的模型需落地为 API 服务。本包提供从 PyTorch → ONNX → TensorRT 的完整链路,并验证端到端延迟。

5.1 ONNX 导出:规避 dynamic axes 陷阱

U-Net 的 skip connection 要求输入尺寸固定,但实际业务中图像尺寸多变。解决方案:导出时指定dynamic_axes,但仅对 batch 维度开放,空间维度固定为 1024×1024

# export_onnx.py dummy_input = torch.randn(1, 3, 1024, 1024).cuda() torch.onnx.export( model, dummy_input, "inpainting.onnx", input_names=["input"], output_names=["output"], dynamic_axes={ "input": {0: "batch_size"}, # 仅 batch 可变 "output": {0: "batch_size"} }, opset_version=12, do_constant_folding=True )

注意:opset_version=12是 TensorRT 8.4 支持的最高版本,opset_version=13会导致 TRT 解析失败。

5.2 TensorRT 引擎构建:INT8 量化与 profile 优化

使用trtexec构建引擎(--int8 --calib=test_data.bin):

  • 校准数据:从 val set 随机采样 500 张图,经normalize后保存为.bin(CHW, float32, row-major)
  • profile 优化:指定--minShapes="input:1x3x1024x1024"--optShapes="input:4x3x1024x1024"--maxShapes="input:8x3x1024x1024",覆盖常见 batch size
trtexec --onnx=inpainting.onnx \ --int8 \ --calib=test_data.bin \ --minShapes="input:1x3x1024x1024" \ --optShapes="input:4x3x1024x1024" \ --maxShapes="input:8x3x1024x1024" \ --workspace=2048 \ --saveEngine=inpainting_int8.trt

5.3 FastAPI 服务封装:异步推理与内存管理

app.py关键设计:

  • 模型单例:全局加载 TRT engine,避免每次请求重建 context
  • 异步队列:asyncio.Queue缓冲请求,防止 burst 请求压垮 GPU
  • 内存释放:每次推理后调用engine.destroy()(TRT 8.4 必须显式释放)
# app.py class TRTInferencer: def __init__(self, engine_path: str): self.engine = self._load_engine(engine_path) # trt.IRuntime.deserialize_cuda_engine(...) self.context = self.engine.create_execution_context() async def infer(self, image: np.ndarray) -> np.ndarray: # 1. copy to device buffer cuda.memcpy_htod_async(self.d_input, image.astype(np.float32), self.stream) # 2. execute self.context.execute_async_v2(bindings=[int(self.d_input), int(self.d_output)], stream_handle=self.stream.handle) # 3. copy back cuda.memcpy_dtoh_async(self.h_output, self.d_output, self.stream) self.stream.synchronize() return self.h_output.reshape(3, 1024, 1024) # 全局实例 inferencer = TRTInferencer("inpainting_int8.trt") @app.post("/inpaint") async def inpaint_endpoint(file: UploadFile = File(...)): image = await file.read() img_array = cv2.imdecode(np.frombuffer(image, np.uint8), cv2.IMREAD_COLOR) # resize to 1024x1024, normalize... result = await inferencer.infer(processed_img) return {"result": base64.b64encode(result.tobytes()).decode()}

提示:trtexec构建的 engine 文件大小约 1.2GB,部署时需确保/tmp有足够空间(默认 2GB),否则trtexec报out of memory错误。


6. 验证不是“看一眼效果图”:PSNR/SSIM/LPIPS 三指标联动分析与业务阈值设定

效果评估不能只看图。本包提供evaluate.py脚本,输出三指标并生成诊断报告,核心是识别“高 PSNR 但低 LPIPS”的伪修复。

6.1 三指标物理意义与业务映射

指标计算方式敏感度业务含义合格阈值(ICVL-IR)
PSNR10*log10(MAX²/MSE)像素级误差修复区域平均保真度≥ 28.5 dB
SSIM结构相似性(亮度/对比度/结构)局部结构边缘与纹理连贯性≥ 0.82
LPIPSVGG 特征空间 L2 距离语义感知人眼主观质量(避免模糊/伪影)≤ 0.15

关键洞察:当 PSNR ≥ 29.0 但 LPIPS > 0.18 时,90% 概率存在高频伪影(如摩尔纹、振铃效应),需检查 generator 的空洞卷积 dilation 值是否过大。

6.2 自动化诊断报告生成

evaluate.py输出report.html,含三类图表:

  • 散点图:PSNR vs LPIPS,标出阈值线(红虚线:PSNR=28.5, LPIPS=0.15)
  • 热力图:各退化类型(scratch/mask/superres)的指标分布
  • Top-5 Bad Cases:按(LPIPS - 0.15) * 1000 + (28.5 - PSNR)排序,定位最差样本
# evaluate.py def calc_lpips(img1, img2, lpips_net): # 使用官方 lpips==0.1.4,预训练 alex net return lpips_net(img1, img2).item() def generate_report(results: List[Dict]): # results[i] = {"psnr": 28.7, "ssim": 0.83, "lpips": 0.12, "deg_type": "scratch"} df = pd.DataFrame(results) # 计算综合得分(越小越好) df["score"] = (df["lpips"] - 0.15) * 1000 + (28.5 - df["psnr"]) df = df.sort_values("score", ascending=False).head(5) # 生成 HTML 报告...

6.3 业务阈值动态校准:基于产线反馈的迭代机制

我们曾遇到 PSNR=29.2 但质检员拒收的情况——原因是修复区域出现0.5px 宽的亮边(人眼敏感,但 PSNR/SSIM 不敏感)。解决方案:

  • 新增边缘锐度检测:用 Sobel 算子提取修复区域边缘,计算std(edge_map),阈值设为 0.08
  • 建立反馈闭环:将质检员标记的“拒收图”加入hard_negative/目录,每周 retrain 时hard_negative_weight=2.0
# train.py 中 hard negative 加权 if sample_path in hard_negative_list: loss = loss * 2.0 # 加重惩罚

从那以后我每次上线新模型,都强制走一遍evaluate.py --report --hard_negative,把报告发给质检组长签字确认——不是因为信不过算法,而是信不过自己没看见的角落。希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询