ARTICLE DETAIL

资讯详情

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

手写试卷擦除实战:基于U-Net图像翻译与ONNX部署

手写试卷擦除实战:基于U-Net图像翻译与ONNX部署 简介这套手写试卷擦除工具是一个基于Python与深度学习的开源项目集成BiSeNetV2、NAFA、SA-GAN等轻量级语义分割和图像修复模型可自动识别并擦除试卷上的手写笔迹同时保留印刷体内容所有模块均已本地验证可运行支持自定义试卷图像输入。项目涵盖模型训练、推理预测、评估测试、数据加载、损失函数、工具函数以及ONNX导出等完整流程代码结构清晰模块职责明确便于理解原理与二次开发。压缩包共69个文件包含45个Python脚本、6个Shell脚本、多个说明与项目文档大小约472KB其中Python脚本覆盖训练、推理、测试与ONNX转换Shell脚本提供一键运行入口说明文档详解环境配置、数据准备、训练命令、推理示例及常见问题解答。已有58人学习下载适合计算机、人工智能、电子信息等专业本科生毕业设计、课程设计或大作业也适合深度学习图像修复方向的入门实战可在此基础上替换主干网络、接入新数据集或优化擦除边缘效果。1. 项目背景与核心思路拆解1.1 为什么要做手写试卷擦除这个工具先说说这个项目到底解决什么问题。教学场景里有个特别常见的需求老师手里有一套好题但已经被学生写过了想再打印出来给下一届用或者想把空白的卷子电子化归档。手动PS擦除太慢而且容易把印刷体也擦花。更麻烦的是试卷扫描件里手写墨迹的颜色深浅、笔迹粗细、覆盖角度各不相同用传统的图像处理算法——比如阈值分割加inpainting——做出来效果普遍不理想手写痕迹一重就残留。这个项目的核心思路是把擦除手写当成一个图像翻译任务来解输入一张含手写的试卷图输出一张干净的空白卷。本质上和超分辨率、去雨、去雾是同一类问题属于像素到像素的深度学习任务。用深度学习模型来做优势在于它不是靠某个固定阈值去判断哪里是手写而是学出了空卷长什么样的分布哪怕手写压住了印刷体的一部分也能把底下的印刷结构合理重建出来。我选用PyTorch作为主力框架配合ONNX做跨平台部署。选择PyTorch的原因很直接生态成熟GitHub上现成的预训练模型多调试也直观ONNX则是因为最终要在Windows、Linux甚至嵌入式设备上跑总不能让每台机器都装一套PyTorch环境。1.2 整体技术方案选型整个项目分为四个环节数据准备、模型训练、测试评估、ONNX转换部署。技术栈和版本我列一下方便你在自己机器上复现时不踩版本坑Python 3.93.8也行但3.9的typing支持更好PyTorch 2.1.x CUDA 11.8如果你显卡驱动不支持CUDA 11.8降到1.13也能跑只是训练速度慢一些OpenCV 4.8图像读写和预处理Albumentations数据增强比torchvision的transforms更顺手ONNX Runtime 1.16导出与推理模型选型上我用的是U-Net作为backbone损失函数走的是L1 Perceptual Loss的组合。为什么不直接上GAN因为GAN训练不稳定对新手不友好而且试卷擦除任务不追求生成多奇幻的风格它追求的是干净、忠实原卷结构L1损失天然合适。后面我还试过加一个判别器做PatchGAN效果提升有限但训练时间涨了一半性价比不高最终版本舍弃了。注意这个项目的核心难点并不在模型结构多花哨而在数据构造和损失函数设计。如果你直接拿公开的街景去模糊数据集来训练模型学到的东西跟试卷擦除完全两回事。2. 数据准备与标注策略2.1 合成数据的构造方法深度学习项目里数据决定了效果上限。手写试卷擦除这个任务现实中很难拿到同一张卷子既有写过的又有空白的成对数据。所以我的做法是——合成。合成思路不复杂拿一批空白试卷扫描图作为背景找一批真实手写笔迹图像作为前景用随机仿射变换、随机颜色扰动、随机透明度混合把笔迹贴到空白卷上生成带手写的样本。这样天然就有一一对应的标签对。具体参数如下背景图至少准备50张以上不同排版、不同扫描亮度的空白卷分辨率统一缩放到640×896保持横纵比在2:3左右太扁或太方的卷子比例会影响模型泛化。手写笔迹来源是公开的手写数据集比如IAM手写数据库的子集以及网上搜集的试卷扫描件。注意笔迹要多样化——签字笔、圆珠笔、铅笔的效果差异很大如果模型只在签字笔上训练遇到铅笔写的卷子基本报废。叠加方式笔迹区域先做高斯模糊核大小随机3~5再做透视变换旋转范围±15度缩放0.9~1.1最后用透明度alpha0.7~0.95叠加到背景上。为了让模型更鲁棒我还会在背景上加高斯噪声和亮度扰动模拟不同扫描仪的底噪。这一步是最耗时间的但数据质量直接决定模型上限。我前前后后折腾了两周反复调整合成参数才让模型在真实扫描件上有可用表现。2.2 数据增强与数据集划分合成数据虽然量大我生成了约8000对但也有个问题太干净。为了抹平合成数据与真实扫描件的分布差异必须做增强。我用的增强组合RandomBrightnessContrast亮度对比度概率0.5幅度0.2GaussNoise噪声概率0.3ShiftScaleRotate平移缩放旋转概率0.4RandomResizedCrop随机裁剪缩放概率0.5尺度范围0.7~1.0训练集、验证集、测试集的划分比例是训练6000对、验证1000对、测试1000对。验证集和测试集要确保包含一部分真实手写扫描件不能全用合成数据否则模型在你自己的测试集上表演完美一上真实战场就翻车。这里有个很关键的心得训练集可以全部合成但验证集和测试集必须掺真数据。哪怕是拿手机拍的带手写的卷子模糊一点没关系能反映真实分布就行。我最初测试集也全是合成数据模型在val loss上表现很好结果拿真实扫描件一测效果惨不忍睹后来换了测试集策略才校正过来。3. 模型训练全流程3.1 网络结构设计与实现模型主体用的是U-Net结构编码器部分我做了轻量化调整没有直接用ResNet34这种大网络而是用了一个6层的卷积编码器每层通道数是[64, 128, 256, 512, 512, 512]解码器对称。为什么不用残差网络因为试卷擦除任务不需要特别大的感受野手写笔画本身的局部性很强——你只需看周围几十个像素就能推断出下面的印刷体长什么样。大网络反而容易在小数据集上过拟合。核心代码结构如下import torch import torch.nn as nn class UNet(nn.Module): def __init__(self, in_channels3, out_channels3): super().__init__() # encoder self.enc1 self._block(in_channels, 64) self.enc2 self._block(64, 128) self.enc3 self._block(128, 256) self.enc4 self._block(256, 512) self.pool nn.MaxPool2d(2) # bottleneck self.bottleneck self._block(512, 512) # decoder self.up4 nn.ConvTranspose2d(512, 512, 2, stride2) self.dec4 self._block(1024, 512) self.up3 nn.ConvTranspose2d(512, 256, 2, stride2) self.dec3 self._block(512, 256) self.up2 nn.ConvTranspose2d(256, 128, 2, stride2) self.dec2 self._block(256, 128) self.up1 nn.ConvTranspose2d(128, 64, 2, stride2) self.dec1 self._block(128, 64) self.out nn.Conv2d(64, out_channels, 1) def _block(self, in_ch, out_ch): return nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue) ) def forward(self, x): e1 self.enc1(x) e2 self.enc2(self.pool(e1)) e3 self.enc3(self.pool(e2)) e4 self.enc4(self.pool(e3)) b self.bottleneck(self.pool(e4)) d4 self.dec4(torch.cat([self.up4(b), e4], dim1)) d3 self.dec3(torch.cat([self.up3(d4), e3], dim1)) d2 self.dec2(torch.cat([self.up2(d3), e2], dim1)) d1 self.dec1(torch.cat([self.up1(d2), e1], dim1)) return torch.sigmoid(self.out(d1))小细节最后一层用了sigmoid而不是纯线性输出因为输入图像做了归一化到[0,1]输出限定在[0,1]区间可以让训练更稳定也能防止输出越界导致图像出现奇怪的色偏。3.2 损失函数与训练参数解析损失函数这块我踩过不少坑这里展开讲讲。早期版本只用了L1 Loss训练出来边缘是糊的。原因是L1 Loss对每个像素一视同仁模型把模糊化当作一个低成本策略——反正模糊了L1也不会高太多。后来我加上了Perceptual Loss感知损失做法是先把生成图像和真实干净图像同时送入一个预训练的VGG16取浅层特征做L1距离。这样模型就不敢随便模糊了因为模糊图像在VGG特征空间的表示会偏离真实图像很远。最终的损失公式Loss 0.6 * L1_loss(output, target) 0.4 * perceptual_loss(output, target)实验下来L1权重0.6、感知权重0.4时在验证集上的PSNR峰值信噪比和SSIM结构相似性综合最优。只加GAN的时候收敛时间多了一倍但PSNR反而降了0.3dB果断放弃。优化器用的是Adam学习率初始1e-4batch size 8显存8GB以下的机器建议降到4训练总轮数120轮。学习率在第60轮和第90轮各衰减一次衰减因子0.1。这样设计是因为前期让模型快速收敛后期用小学习率精细打磨细节。3.3 训练过程监控与调参经验训练日志我每5轮打印一次记录train loss、val loss、PSNR三个指标。PSNR计算方式很简单import numpy as np def calculate_psnr(img1, img2, max_val1.0): mse np.mean((img1 - img2) ** 2) if mse 0: return float(inf) return 20 * np.log10(max_val / np.sqrt(mse))实践经验val loss从第40轮开始下降明显变慢第80轮之后基本持平PSNR在第80轮左右达到峰值之后训练损失还在降但val指标开始波动这就是过拟合信号建议早停。此外训练过程中记得开启梯度裁剪gradient clippingmax_norm设为1.0。这个任务输入输出的数值范围比较窄[0,1]区间但其实梯度爆炸的概率并不低尤其是batch size小、使用了BatchNorm的情况下。我中途遇到过loss突然变成NaN的情况排查了半天才发现是某几张图像里有纯黑区域配合高学习率导致梯度爆炸加上梯度裁剪之后训练就稳定了。4. 测试环节与效果评估4.1 测试评估指标与可视化验证测试阶段不能只盯着PSNR和SSIM这两个指标有个通病分数高不代表人眼看着舒服。PSNR对全局像素误差敏感但它会把输出整体变灰误判为低误差。所以我额外制作了一个残差热力图把输出图像和干净标签做逐像素差差值映射到color map上可视化。这样一眼就能看出模型是不是把印刷体也擦掉了一块——有这种问题的话热力图在印刷体边缘区域会特别亮。测试时数据加载要注意图像不能直接resize就喂给模型不然长宽比变了印刷体文字会被拉伸。我在测试管线里加了等比例resize 边缘padding的逻辑把图像统一处理成640×896保持内容不变形。这个细节直接影响最终效果尤其是卷头标题区域的文字结构。测试集上我统计了几个关键数字指标合成测试集真实扫描测试集PSNR34.2 dB31.8 dBSSIM0.9670.941平均推理耗时(单张)45ms(GPU)45ms(GPU)真实测试集比合成测试集低2.4dB这个差距在预期内但肉眼观感仍然干净利落。如果你的项目要求更严格建议扩充真实数据到整个训练集的20%以上再配合弱监督策略来做域适应。4.2 失败案例分析测试过程中我专门挑了一些刁钻样本出来看铅笔写痕铅笔灰度低、与印刷体灰度接近模型有时候会误判把印刷体的一部分也擦掉。优化方案在合成阶段把一部分前景笔画调低透明度模拟铅笔效果。手写压线下划线、表格线当手写和印刷线条重合时模型倾向于把整条线擦掉。这种问题靠数据增强无法完全解决只能接受一定程度的线条断裂或者在后处理阶段用形态学闭运算把断裂的线条修复回来。红色笔迹批改红色在RGB空间和黑色墨迹差异很大如果训练数据里没有或很少红色笔迹模型看到红色基本不处理。后来我在合成数据里专门加了一批红色笔迹图像算是解决了。这些失败案例给我的启示是深度学习模型不会像人一样智能地绕过印刷体它只是学到了统计规律。要提升特定场景的效果唯一的办法是让那个场景出现在训练数据里。5. ONNX转换全流程与部署实践5.1 PyTorch模型导出ONNX的正确姿势模型训练完成并验证合格后就到了部署环节。ONNX转换这一步看似简单但实际操作中有不少细节。我先说最基础的导出代码import torch import onnx import onnxruntime as ort # 加载训练好的权重 model UNet(in_channels3, out_channels3) model.load_state_dict(torch.load(best_model.pth)) model.eval() # 构造输入张量 dummy_input torch.randn(1, 3, 640, 896) # 导出ONNX torch.onnx.export( model, dummy_input, eraser.onnx, export_paramsTrue, opset_version12, do_constant_foldingTrue, input_names[input], output_names[output], dynamic_axes{ input: {0: batch_size}, output: {0: batch_size} } )几个值得注意的点opset_version我用的是12兼容性最好。版本太老不支持某些算子太新对部署环境要求高。dynamic_axes把batch维度设成动态这样导出后的模型既能一次处理1张也能批量处理多张。如果不设onnx模型会绑定死batch1。导出前务必调用model.eval()否则Dropout和BatchNorm的推理行为不对导出的模型推理结果会和PyTorch不一致。5.2 验证ONNX模型输出一致性导出ONNX之后最重要的一步是验证PyTorch模型的输出和ONNX Runtime的输出要基本一致。这一步我能理解很多人的心情——觉得自己模型都训好了导出还能出什么问题但实际情况是常常会有1e-4量级的微小差异严重时甚至会有几个像素的明显偏差。验证代码import numpy as np import onnxruntime as ort import torch # 用同一张测试图 test_img torch.randn(1, 3, 640, 896) # PyTorch推理 with torch.no_grad(): torch_output model(test_img) # ONNX Runtime推理 ort_session ort.InferenceSession(eraser.onnx) ort_inputs {ort_session.get_inputs()[0].name: test_img.numpy()} ort_output ort_session.run(None, ort_inputs)[0] # 计算最大绝对误差 max_diff np.abs(torch_output.numpy() - ort_output).max() print(fMax diff: {max_diff:.6f}) # 如果diff非常小通常1e-5说明转换成功 assert max_diff 1e-4, ONNX output mismatch!我实际操作中遇到过一次最大差异达到0.02的情况排查后发现是模型里一个nn.Upsample的align_corners参数问题——PyTorch默认align_cornersFalse但ONNX导出时这个参数如果没被正确记录推理结果就会不同。解决办法是把Upsample层显式写成nn.Upsample(scale_factor2, modebilinear, align_cornersFalse)并且在导出前用torch.onnx.export的operator_export_type参数固定算子行为。5.3 ONNX Runtime推理代码模板转换验证通过后部署推理就简单了。我这里给出一份完整的ONNX Runtime推理模板支持批量处理目录下所有图片import onnxruntime as ort import cv2 import numpy as np from pathlib import Path class EraserEngine: def __init__(self, onnx_path, input_size(640, 896)): self.session ort.InferenceSession( onnx_path, providers[CUDAExecutionProvider, CPUExecutionProvider] ) self.input_name self.session.get_inputs()[0].name self.input_size input_size def preprocess(self, img_bgr): 等比例缩放padding到目标尺寸保持内容不变形 h, w img_bgr.shape[:2] th, tw self.input_size scale min(th / h, tw / w) nh, nw int(h * scale), int(w * scale) img cv2.resize(img_bgr, (nw, nh)) # padding到目标尺寸 pad_top (th - nh) // 2 pad_bottom th - nh - pad_top pad_left (tw - nw) // 2 pad_right tw - nw - pad_left img cv2.copyMakeBorder( img, pad_top, pad_bottom, pad_left, pad_right, cv2.BORDER_CONSTANT, value(255, 255, 255) ) img img.astype(np.float32) / 255.0 img img.transpose(2, 0, 1) return img[np.newaxis, ...], (scale, pad_left, pad_top) def postprocess(self, output, meta): 还原到原始尺寸 scale, pad_left, pad_top meta out output[0].transpose(1, 2, 0) out (out * 255).clip(0, 255).astype(np.uint8) if pad_left 0 or pad_top 0: h, w out.shape[:2] crop out[pad_top:h - (pad_top (out.shape[0] - int(self.input_size[0]))), pad_left:w - (pad_left (out.shape[1] - int(self.input_size[1])))] # 上面这行太绕用下面这种更清晰的方式 real_h int(self.input_size[0] / scale) real_w int(self.input_size[1] / scale) crop out[pad_top:pad_top real_h, pad_left:pad_left real_w] return crop def predict(self, img_bgr): meta None input_tensor, meta self.preprocess(img_bgr) output self.session.run(None, {self.input_name: input_tensor})[0] return self.postprocess(output, meta) # 使用示例 engine EraserEngine(eraser.onnx) img cv2.imread(test_written.jpg) clean engine.predict(img) cv2.imwrite(test_clean.jpg, clean)这段代码里我故意保留了注释掉的老写法说明一下如果你直接用self.input_size[0] / scale来做后处理裁切更简单也更稳用out.shape反推容易因为padding边界问题多裁或少裁一两个像素。5.4 踩过的部署坑CUDAs加速与动态尺寸ONNX Runtime在GPU上跑大多数时候确实比PyTorch更快因为图优化做得好。但如果你部署的目标机器只有CPU推理一张640×896的图大约需要300~500ms勉强可用但不够流畅。如果想在CPU上提速有几个方向第一转成INT8量化模型体积缩小4倍速度提升2~3倍精度损失大概在1~2dB PSNR左右对于试卷擦除这个任务完全够用。第二换成MobileNet作为编码器骨干模型参数量从30M降到5M以下精度损失约0.5dB但CPU推理时间能压到100ms以内。另外ONNX Runtime的providers参数要按顺序写优先CUDA其次CPU。如果不写全遇到没有GPU的环境时整个程序会直接报错而不是自动fallback到CPU。这算是部署新手最常见的坑之一。6. 优化方向与扩展思路整个项目做到这个程度已经可以在本地顺畅跑通输入含手写的试卷图 - 输出干净空白卷的完整流程了。如果要进一步优化我会建议按以下优先级推进第一升级为边缘引导的生成网络。在U-Net输入侧并联一个边缘检测分支用Sobel或Canny提取印刷体边缘让模型显式地看到哪些结构必须保留。第二引入对比学习做域自适应。合成数据和真实扫描件的风格差异是模型泛化能力的天花板。用一个预训练的特征提取器把真实无标注数据的特征拉近合成数据的特征空间这一招能再提2~3dB PSNR。第三做成Web端轻量应用。ONNX模型可以直接用ONNX Runtime Web在浏览器里跑把整个项目打包成一个HTML JS的页面上传试卷图就出结果不需要任何后端服务。我用WebAssembly版ONNX Runtime试过推理速度在普通笔记本上约200ms体验很好。我在实际测试中还发现一个有趣的现象模型对印刷体手写的混合输入效果很好但如果直接输入纯手写笔记比如一张白纸上的手写内容模型会把所有内容都擦掉输出一张纯白图。这其实暴露了模型的本质——它学到的是一种背景重建能力而不是识别手写并分离的能力。理解这个边界很重要能帮你合理设定工具的使用范围不过度期待。复盘整个项目最有价值的经验谈不上模型多精妙反而是数据构造和损失函数微调这两个环节投入产出比最高。训练网络本身反而不怎么费心PyTorch生态把这些都简化了。如果你要复现这个项目建议把时间重点放在数据多样性上——手写风格、扫描清晰度、笔迹颜色这三样做得越丰富模型的鲁棒性就越强其他都是锦上添花。本文还有配套的精品资源点击获取
返回列表