两阶段图像修复:基于GAN的试卷手写擦除技术解析
2026/9/17 4:52:08 网站建设 项目流程

简介:面向教育图像处理与OCR预处理场景,提供一套基于深度学习的试卷手写文字擦除完整实现。资源包含两阶段训练策略,先以dice_loss与l1 loss联合优化,再仅用l1 loss微调,并通过随机crop、横向翻转与小角度旋转增强数据;测试阶段采用分块与交错分块预测,配合镜像padding和横向增强,融合多模型结果以提升擦除稳定性。压缩包共30个文件,以Python脚本为主(22个py),另有3个shell脚本、2个readme、1个txt说明、1个docx手册及1个子zip模型文件,整体仅823KB,结构清晰,适合有一定深度学习基础的研究者或开发者直接阅读和迁移。已有890人学习下载,可作为试卷去手写、水印智能消除等场景的参考基线,亦可用于复现论文方法或二次开发。

1. 试卷手写擦除为什么比想象中难

试卷擦除的输入是一张扫描图:印刷体题目、手写答案、批改笔迹叠在同一平面上。大多数图像修复模型能抹掉大块污渍,但碰到手写笔迹与印刷体笔画交叉的地方,经常把印刷体也一起擦掉,留下断笔和残影。这个项目没有走端到端inpainting的路线,而是把任务拆成两步:先让网络学会定位手写mask,再放下mask监督、专心做重建。拿到的源码里带SA-GAN、NAFA、BiSeNetV2三套网络实现,训练用512×512随机patch,测试用交错分块加镜像padding,最后融合两个模型的输出。这套链路对扫描文档去手写、旧试卷清空、表单去水印这类场景是可以直接复用的,对正在做文档图像预处理、OCR前处理或图像修复的工程师,有不少可以抄走的细节。

2. 网络骨架与损失设计:自注意力GAN、非局部块与两阶段切换

2.1 从文件命名看训练主链路

解压后第一眼就能看到train.py、Model.py、sa_gan.py、nafa_archv1.py、BiSeNetV2.py、non_local.py、discriminator.py这一组文件。train.py是训练入口,Model.py负责把生成器、判别器、网络分支组装起来;sa_gan.py是带自注意力的GAN生成器,non_local.py提供非局部注意力块,BiSeNetV2.py是双向分割网络的结构实现,nafa_archv1.py则是另一套面向图像重建的特征聚合架构。从文件组织来看,项目实际是并行训练两套生成器,测试阶段再对两套模型的输出做融合,而不是二选一。

这种做法的思路是让不同归纳偏置的网络互相纠错。自注意力生成器的感受野大,对文字的笔画走向和长距离上下文更敏感;NAFA这类特征聚合结构对局部边缘形态的保持更强。两个模型在同一位置的错误通常不会重叠,平均之后能明显压掉单模型的病态输出。non_local.py在这套结构里的角色是给生成器中间特征图加全局依赖,自注意力块的计算量随特征图尺寸平方增长,所以一般只挂在分辨率较低的特征层上,分辨率高的层用普通卷积保留细节,属于典型的“大感受野加局部精度”搭配。

2.2 生成器与判别器的配合方式

discriminator.py实现了判别器,但GAN损失在这个项目里不是主力。从损失文件分布可以看出,主导训练的是dice_loss和L1 loss,对抗损失只是辅助信号,用来让重建图像在感知上更自然。常见做法是给对抗损失一个很小的权重,比如L1权重设为10、GAN权重设为0.1到0.2。如果权重给大了,训练曲线会反复震荡,PSNR不涨反跌,因为判别器把注意力放在笔迹纹理的真假上,而擦除任务的首要目标是像素级地接近干净底版。

我一般会把它理解成一个带先验约束的修复问题:L1负责数值准确性,dice负责区域定位,GAN负责让结果看起来像扫描件而不是磨皮后的塑料。三者各司其职,任何一个权重失衡都会在验证集上暴露出来,通常表现为mask区域糊掉或边缘出现条状伪影。训练日志里如果只看PSNR,GAN分支的贡献不大,但对人眼观感影响明显,尤其是手写笔迹那种密集细长的笔画,没有判别器的时候容易糊成一片。

2.3 两阶段损失的具体实现

第一阶段用dice_loss加L1 loss,第二阶段只保留L1 loss,这个切换是项目的核心。dice_loss在训练前期能把掩码轮廓快速拉到位,因为它直接优化掩码与标定的重合度,对大块笔迹尤其高效。但训练到后期,同batch里不同尺度笔迹的掩码尺寸可能差一个数量级,dice梯度会偏向大块区域,小块笔迹的预测逐步退化,继续使用反而拖累重建质量。此时切到纯L1 loss,让所有像素点以同样的权重参与回归,数字上更稳。

# losses.py 两阶段损失切换的核心逻辑 class StageLoss(nn.Module): def __init__(self, stage='phase1'): super().__init__() self.stage = stage self.l1 = nn.L1Loss() def dice_loss(self, pred_mask, gt_mask): pred = pred_mask.sigmoid().flatten(1) gt = gt_mask.flatten(1) intersection = (pred * gt).sum(1) union = pred.sum(1) + gt.sum(1) return 1 - (2 * intersection + 1e-6) / (union + 1e-6) def forward(self, out, batch): if self.stage == 'phase1': return 0.7 * self.dice_loss(out['mask'], batch['mask']) \ + self.l1(out['image'], batch['clean']) return self.l1(out['image'], batch['clean'])

flatten(1)把每个batch样本的掩码拉成一维,在单样本内部算交并比,再做batch平均,避免掩码大小不同的样本互相影响;union加1e-6是防空掩码除零。权重0.7可以按验证集表现调整,第一阶段结束时要观察mask的IOU是否稳定在目标区间,再进入第二阶段。阶段切换的时机,我一般会在验证集PSNR出现平台期之后再切,具体轮数要看训练曲线,常见做法是阶段二在总epoch的60%到70%位置开启。

阶段主损失监督信号训练目标
phase1dice_loss + L1mask + 重建图先把该擦的位置找对
phase2L1 only重建图把擦除后的像素修补自然

2.4 EMA与checkpoint转换

train.py里还带ema.py,这组滑动平均权重在训练过程中持续累积参数的指数平均。推理时直接使用EMA权重,最终效果通常比最后一次step的权重PSNR高0.2到0.4dB,因为EMA等效于把训练后期参数的高频抖动磨平了。ckpt_convert.py的任务就是把训练状态文件里带ema_前缀的键值剥出来,重新组装成Model.py能直接加载的state_dict。

提示:如果训练时开了DataParallel,权重键名会带module.前缀,ckpt_convert.py里记得统一去掉,否则加载时会报key不匹配。

3. 数据管线:横向翻转、小角度旋转与512×512随机裁剪

3.1 配对数据的组织方式

data/dataloader.py按对读取数据:带手写的扫描图和对应的干净底版。训练时L1 loss需要干净底版做回归目标,所以样本必须配对。compute_mask.py负责生成掩码,常见做法是先对两张图做逐像素差分,再经过阈值化和形态学膨胀得到手写区域。这里有个容易翻车的细节:扫描仪的光学特性会让两次扫描同一张纸时灰阶也不完全一致,差分图里印刷体边缘也会产生响应。膨胀核尺寸开大了,掩码就会把印刷体笔画包进去,训练出来的模型会连印刷体一起抹掉。我一般会把膨胀控制在3到5个像素,并观察验证集上印刷体的保留情况。

3.2 增强策略为什么克制

只用了横向翻转和小角度旋转,没加颜色抖动、随机亮度、大幅旋转,这是刻意为之。印刷体文字的方向先验对擦除结果影响很大:字符语义依赖朝向,旋转超过一定角度会破坏网络对文字结构的理解。颜色抖动则会让L1损失在RGB通道上失真,因为擦除任务关心的是结构差异,不是光照变化下的鲁棒性。横向翻转对文字识别是安全的,网络对镜像后的字符仍然能保持结构响应。

增强项参数范围对文字先验的影响是否使用
横向翻转p=0.5无破坏使用
小角度旋转±5°以内基本无影响使用
大角度旋转超过15°破坏字符朝向语义不用
颜色抖动任意范围干扰L1重建不用

3.3 随机crop 512×512 patch训练

扫描图原始尺寸通常是3000×2000甚至更大,整图送进网络显存装不下,一张图也只是一个样本。随机裁剪成512×512的patch后,一张图能产生几十个训练样本,batch维度也能撑起来。crop的关键是img、clean、mask三张图必须用同一组偏移,否则损失函数在两个错位的输入上计算,训练永远不会收敛。

def random_crop_pair(img, clean, mask, crop=512): # 同一组偏移裁剪三张图,保证mask与图像语义对齐 _, H, W = img.shape y = np.random.randint(0, H - crop + 1) x = np.random.randint(0, W - crop + 1) return (img[:, y:y + crop, x:x + crop], clean[:, y:y + crop, x:x + crop], mask[:, y:y + crop, x:x + crop])

参数crop设为512时,如果某张图像短边不足512,预处理阶段就要先做镜像padding到安全尺寸再裁剪。补零会产生亮度跳变,擦除结果的边缘容易带黑边,不建议用。训练时分两个阶段跑,train.sh里可以这样串起来:

# 阶段一:让mask先收敛,lr给大一点 python train.py --model sa_gan --phase 1 \ --dataroot data/dehw --crop 512 --batch_size 8 \ --lr 1e-4 --epoch 60 # 阶段二:加载阶段一权重,只优化L1重建 python train.py --model sa_gan --phase 2 \ --ckpt logs/phase1/best.pth --crop 512 \ --batch_size 8 --lr 5e-5 --epoch 40

batch_size 8配合512×512输入,在16G显存的卡上差不多是上限,如果显存不够可以把batch_size降到4,同时适当调低学习率,避免梯度更新步长过大。phase1的epoch数可以看mask IOU曲线来定,验证集IOU连续若干轮不涨时就可以切phase2。

3.4 高斯模糊让掩码边缘平滑

gauss.py实现了高斯核来平滑mask。直接由差分图二值化得到的mask边缘是硬边界,dice_loss在硬边界处的梯度方向不稳定,训练后期容易抖动。对mask做高斯模糊后,置信度从0渐变到1,L1 loss在边缘的权重也平滑过渡,重建阶段对边缘像素的处理会更自然。这里的sigma取值不需要大,2到3个像素的模糊量就足够把边界软化。

4. 推理策略:交错分块、镜像padding与双模型融合

4.1 全图推理的两个问题

直接把整张扫描图resize到模型输入尺寸,每个patch的语义都被压缩,印刷体的高频笔画细节会丢失,输出看起来是“虚”的。另一个问题是网络感受野固定,原图的文字尺度被压扁后,卷积核分不清笔画边缘和背景噪声。测试阶段必须保留原分辨率,在图上滑动窗口分块预测,再用某种方式把分块结果拼回整图。项目源码里的predict.py、test.sh都是围绕这个逻辑组织的。

4.2 镜像padding与保留中心预测

朴素分块的问题是窗口边缘缺乏上下文。卷积核在边缘拿不到完整的邻近像素,预测结果在块边界区域会差很多。摘要里提到的做法是先给测试图做镜像padding,然后让分块与分块之间保留重叠区域,每次只取预测结果的中心部分。镜像padding用反射系数填充,比补零柔和,对图像修复类任务几乎是标配操作。重叠与保留中心的组合,让每块预测都使用“看到更多上下文”的窗口,输出时只信任中心区域的结果。

4.3 交错分块:偏移起点把边界变成中心

即使每块只保留中心,窗口起点固定时边界仍然会落在整图固定的位置,这些位置正好是两次预测的过渡区。交错分块的做法是把整体起始坐标再偏移一段距离,重新跑一遍分块推理,让第一轮的窗口边界落在第二轮的中心区域,融合后网格感被明显打散。

def overlap_predict(model, img, patch=512, keep=256, offset=128, device='cuda'): # img: [1,3,H,W], 保持原分辨率,重叠区域加权平均 import torch.nn.functional as F _, _, H, W = img.shape pad_h = patch - H % patch if H % patch else 0 pad_w = patch - W % patch if W % patch else 0 img = F.pad(img, (pad_w // 2, pad_w - pad_w // 2, pad_h // 2, pad_h - pad_h // 2), mode='reflect') _, _, Ph, Pw = img.shape acc = torch.zeros((1, 3, Ph, Pw), device=device) cnt = torch.zeros((1, 3, Ph, Pw), device=device) step = patch - keep y0 = x0 = (patch - keep) // 2 for off in (0, offset): for y in range(off, Ph - patch + 1, step): for x in range(off, Pw - patch + 1, step): with torch.no_grad(): pred = model(img[:, :, y:y + patch, x:x + patch]) acc[:, :, y + y0:y + y0 + keep, x + x0:x + x0 + keep] += \ pred[:, :, y0:y0 + keep, x0:x0 + keep] cnt[:, :, y + y0:y + y0 + keep, x + x0:x + x0 + keep] += 1 return (acc / cnt.clamp(min=1))[:, :, :H, :W]

step取patch减keep,也就是256,第一轮各窗口的中心区域刚好首尾拼接;第二轮偏移offset,中心区域与第一轮有重叠,cnt在这些位置等于2,最后除以cnt得到加权平均。keep越大,单块保留区域越多、计算越快,但块边缘的伪影也越明显;offset取patch的四分之一到一半之间效果差别不大,128是我常用的值。这里的镜像padding要求pad尺寸不能超过原图对应边长,如果图特别小,需要先把图整体resize到比patch大再进入这个函数,否则reflect模式会直接报错。

4.4 横向镜像TTA与双模型融合

测试时对输入做横向翻转,预测完再把输出翻回来与原结果平均,这是TTA的标准用法。翻转不会改变文字语义,但网络对翻转前后同一位置的响应会有细微差异,平均后能压制部分噪声。双模型融合是在此基础上再做一层平均:两个生成器各自过一遍TTA,最终结果取两个预测的平均。权重不必学,验证集上观察哪个模型更强,按0.6/0.4或0.5/0.5手动配平即可。

bash test.sh --input_dir ./inputs --output_dir ./outputs \ --patch 512 --keep 256 --offset 128 --tta True --ensemble True
推理配置边缘伪影网格感相对耗时
全图resize高(结构失真)1x
直接分块明显明显1.2x
分块+保留中心较弱较弱2x
交错分块+保留中心基本消除4x
交错+双模型+TTA最弱消除8x

5. ONNX导出与部署验证:固定输入尺寸与PSNR检查

5.1 先把EMA权重还原成推理权重

convert_onnx.py导出前要先把训练权重整理干净。直接用训练产物加载会遇到两个常见问题:EMA字段的键名不匹配,以及DataParallel带来的module.前缀。我的习惯是先跑一遍ckpt_convert.py生成inference.pth,再执行导出,不要在导出脚本里反复处理键名。

python ckpt_convert.py --ckpt logs/phase2/best_ema.pth --out weights/inference.pth python convert_onnx.py --ckpt weights/inference.pth --out dehw.onnx --size 512

ckpt_convert只做键名重映射和数据拷贝,不涉及计算图操作,所以CPU上就能完成,不需要GPU。

5.2 导出时动态轴取舍

ONNX导出前要决定输入shape。动态H、W会带来部署端的算子兼容问题,固定512×512则可以让推理后端做更激进的图优化。实际部署时推理逻辑本来就按512的patch滑动,完全没有动态分辨率的需求,所以我只对batch轴保留动态维度,H、W固定不动。

import torch from nafa_archv1 import NAFANet model = NAFANet(out_channels=3) # 按训练时配置实例化 model.load_state_dict(torch.load('weights/inference.pth', map_location='cpu')) model.eval() torch.onnx.export( model, torch.randn(1, 3, 512, 512), 'dehw.onnx', input_names=['input'], output_names=['output'], dynamic_axes={'input': {0: 'batch'}, 'output': {0: 'batch'}}, opset_version=11 )

dummy输入是1×3×512×512,等于推理时最常用的batch为1的patch;opset_version=11的兼容性最好,如果后端支持可以升到13或15,部分注意力算子在更高opset下会生成更简洁的图。导出后建议用onnxruntime跑同一个patch,对比PyTorch输出的最大绝对误差,通常应小于1e-4。

5.3 用PSNR与差值图检查部署效果

PSNRLoss.py在训练时监控重建图与干净底的峰值信噪比,但部署验证不能只看全图PSNR。印刷体笔画如果被擦掉一半,PSNR会掉得明显;如果只是被轻微模糊,PSNR变化不大,人眼却一眼能看出区别。所以我在验收时固定一组测试图,分别记录mask区域和整图的PSNR,再抽查三处手写笔迹与印刷体交叉的局部截图。一个好用的技巧是把“整图resize预测”和“交错分块预测”的输出做差,差值图上如果印刷体区域有明显的结构响应,说明分块推理确实保住了高频细节,这个信号比单个PSNR数字可靠得多。

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

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

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

立即咨询