ARTICLE DETAIL

资讯详情

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

基于深度学习的试卷手写擦除:两阶段训练与分块推理实战

基于深度学习的试卷手写擦除:两阶段训练与分块推理实战 简介这份资源面向深度学习与图像处理方向的学习者和开发者提供一套基于深度学习的试卷手写文字擦除完整实现方案可用于试卷还原、文档图像净化等场景适合具备一定PyTorch基础、希望深入理解图像修复与生成模型的中高级读者。压缩包共30个文件约94KB以22个Python源码为主辅以3个Shell脚本、2个readme及说明文档涵盖数据加载、损失函数、网络模型、mask生成、训练与测试等模块并附模型文件与ckpt转换、ONNX导出等工具脚本。训练采用横向翻转与小角度旋转增强随机裁剪512×512 patch分两阶段优化先以dice_loss加l1 loss再仅保留l1 loss。测试环节使用分块与交错分块策略配合镜像padding和横向镜像增强并融合两个模型的预测结果以提升边缘区域效果。目前已有1397人学习读者可据此复现完整流程掌握数据增强、损失设计、分块推理与模型融合等关键技巧。1. 试卷手写擦除这套源码到底能不能直接跑起来改过学生试卷电子版的老师或做教育信息化的工程师大概率都遇到过同一个麻烦想把学生手写答案从卷面上抹掉只留印刷题干用 PS 的污点修复画笔一张张涂涂到第十张手就废了。这套「基于深度学习实现试卷手写文字擦除」的资源包干的就是把这件事自动化——输入一张带手写笔迹的试卷图输出一张只剩印刷体的干净底图。它属于图像到图像的翻译任务主干是 GAN 加注意力机制的组合配套了训练脚本、测试脚本、损失函数定义、模型文件和一封说明文档。适合两类人一类是想直接拿模型跑推理、批量处理试卷的从业者另一类是拿它当深度学习图像修复练手项目、想拆开看两阶段训练怎么设计的同学。下面按「资源是什么 → 怎么用 → 坑在哪」的顺序把我拆包和复现时踩过的细节讲清楚。2. 拆开压缩包目录结构与两阶段训练的设计逻辑2.1 从文件清单反推工程结构先把包里的文件按职责归一下类这样后面找入口不会乱。资源包解压后大致是这么几块目录/文件职责data/dataloader.py数据加载负责读图、增强、切 patchloss/Loss.py、PSNRLoss.py、losses.py损失函数定义含 dice、l1、PSNR 相关models/sa_gan.py、non_local.py、sa_aidr.py、networks.py、idr.py、Model.py、discriminator.py、BiSeNetV2.py、nafa_archv1.py生成器、判别器、注意力模块、分割骨干compute_mask.py生成手写区域的 mask 文件train.py/train.sh训练入口与启动脚本test.py/test.sh测试入口与启动脚本convert_onnx.py/ckpt_convert.py/ema.py模型导出、权重转换、指数滑动平均utils.py、gauss.py通用工具与高斯相关处理项目说明.md、说明文档.txt、readme使用说明看到sa_gan.py和non_local.py基本能判断生成器里用了自注意力self-attention加非局部块这是擦除类任务里保留全局结构一致性的常见做法。BiSeNetV2.py的出现说明 mask 生成或辅助分支可能借用了轻量分割网络用来定位手写笔迹区域。2.2 两阶段训练为什么这么设计说明文档里写得很明确训练分两阶段第一阶段损失是dice_loss l1 loss第二阶段只保留l1 loss。这个设计不是拍脑袋背后有它的道理。第一阶段加 dice loss本质是让网络先把「哪里是手写、哪里要擦」这个区域判断学准。dice 系数衡量的是预测区域和真实 mask 的重叠度对前景背景极不平衡的场景手写笔迹在整张卷面里占比很小特别友好能逼着网络关注到稀疏的笔迹像素。这个阶段网络学的是「定位」。第二阶段砍掉 dice只留 l1是让网络把精力从「找区域」转到「补像素」。l1 直接约束生成图和干净底图逐像素的差距配合 GAN 的对抗损失让擦除后的区域纹理、纸张底色、印刷体边缘更自然。如果第二阶段还留着 dice网络会过度关注 mask 边界反而在填充内容上偷懒出现擦除区域发灰、和周围纸张对不上的问题。提示两阶段的切换点通常靠一个 epoch 阈值或手动改配置控制具体数值以train.py里的参数为准不同数据集收敛速度不一样别照搬。2.3 数据增强只做翻转和小角度旋转文档里强调增强「仅使用横向翻转和小角度旋转保留文字的先验」。这点值得单独说。很多做图像修复的人习惯性堆一堆增强——随机裁剪、色彩抖动、大角度旋转全上结果在这个任务上翻车。原因是试卷有强先验文字是横排的行有方向印刷体有固定朝向。你要是给它来个 90 度旋转或者垂直翻转网络学到的「文字应该长这样」的先验就被破坏了擦除时容易把印刷体也当成噪声抹掉。横向翻转是安全的因为左右镜像后文字依然可读、行方向不变。小角度旋转一般控制在正负几度模拟的是扫描时的轻微倾斜也在合理范围内。随机 crop 成 512x512 的 patch 训练是为了控制显存同时增加样本多样性这个尺寸后面测试时还要对齐是个关键参数。3. 跑通推理从 mask 生成到 test.sh 的完整链路3.1 环境与依赖的常见配置这套代码是 PyTorch 系依赖无非是 torch、torchvision、opencv、numpy、Pillow 这几样。我一般会先建个干净环境再装避免和系统里的老版本打架conda create -n dehw python3.8 -y conda activate dehw pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install opencv-python numpy pillow tqdm逻辑说明Python 3.8 是这类老项目的稳妥选择太新的版本容易在旧版 torch 上出兼容问题。torch 的 CUDA 版本按你显卡驱动来选cu118 只是示例装之前先用nvidia-smi看驱动支持的 CUDA 上限。opencv 用来读写图和做分块tqdm 是训练时的进度条缺了会在 import 阶段就报错。参数说明--index-url指定 PyTorch 官方源国内直连慢的话换成对应镜像即可但别混用多个源容易装出半残的包。3.2 生成 mask 文件compute_mask.py的作用是给训练/测试数据生成手写区域的标注 mask。这一步是整条链路的前置mask 不对后面全白搭。python compute_mask.py \ --data_dir ./dataset/train \ --mask_dir ./dataset/train_mask \ --img_size 512逻辑说明脚本遍历data_dir下的原图对每张图计算手写笔迹的二值 mask写到mask_dir。mask 里白色255代表要擦除的手写区域黑色0代表保留的印刷体背景。参数说明--data_dir是原始试卷图目录--mask_dir是输出目录得提前建好。--img_size要和训练时的 patch 尺寸一致这里填 512。如果你的 mask 是人工标注的这一步可以跳过直接把标注好的 mask 放进对应目录即可。跑完记得抽查几张确认 mask 边缘没有把印刷体也圈进去。3.3 训练启动与两阶段切换训练入口是train.sh里面封装了train.py的调用。文档说「运行 sh train.sh 生成 mask 并开始训练」说明脚本里可能串了 mask 生成和训练两步。bash train.sh逻辑说明脚本内部一般会先调compute_mask.py再进train.py的主循环。第一阶段用dice_loss l1跑到设定轮数后切第二阶段只留l1。参数说明真正要调的是train.py里的几个关键项——batch_size512 patch 下显存吃紧就降到 4 或 2、lr学习率两阶段切换时通常会衰减、epochs总轮数、阶段切换的 epoch 阈值。这些值脚本里给了默认但换数据集后大概率要重调。训练日志里重点盯两个数第一阶段看 dice 是否稳定下降第二阶段看 l1 是否还在缓慢降如果 l1 早早平了说明要么数据不够要么学习率太小。3.4 测试脚本与分块推理测试是这套代码里最讲究的部分文档列了四条 trick我逐条对应到test.py的行为上讲。bash test.sh逻辑说明test.sh调用test.py加载训练好的权重对测试图做分块预测再拼回整图。文档里的四条 trick 分别是分块测试切 512x512 保持和训练一致、交错分块边缘重复、只保留中心、横向镜像增强、双模型融合。参数说明分块尺寸必须等于训练的 patch 尺寸 512否则分布不一致擦除效果会明显变差。交错分块里的「重复部分」宽度是个可调参数重复越多拼接越平滑但耗时越长一般取块尺寸的 1/4 到 1/2。双模型融合需要你准备两个权重文件在test.py里指定两个模型路径输出取平均或加权平均。4. 避坑与排查这几处不注意跑出来的图没法看4.1 现象擦除区域发灰、和周围纸张对不上原因第二阶段损失里还残留 dice或者第二阶段训练轮数不够网络只顾着圈区域没学会补像素。也可能是 l1 权重设得太小对抗损失压过了重建损失。解决确认第二阶段损失只剩 l1把 l1 的权重适当调大多跑几轮第二阶段。如果还是发灰检查训练数据里干净底图和带手写图的配准是否严格对齐错位一两个像素就会导致填充颜色偏移。4.2 现象印刷体被误擦题干缺字原因数据增强用了大角度旋转或垂直翻转破坏了文字方向先验或者 mask 标注时把印刷体边缘圈进了手写区域。解决把增强严格限制在横向翻转和小角度旋转角度阈值调小。回头抽查 mask把误圈印刷体的样本挑出来重标。另外第一阶段 dice 权重过高也会让网络过度激进地扩大擦除范围适当降一点。4.3 现象分块拼接处有明显接缝原因分块时没有做边缘重复或者重复区域太窄每块预测的边缘质量差直接暴露在拼接线上。解决开启交错分块让相邻块有重叠且只保留每块预测结果的中心部分。重复宽度调大到块尺寸的 1/4 以上。这一步是文档里明确点出的 trick别图省事关掉。4.4 现象单模型效果不稳同一张图时好时坏原因单个模型对某些笔迹风格泛化不够尤其是训练集没覆盖到的字迹。解决用双模型融合把两个不同阶段或不同初始化的权重预测结果做平均。文档里「测试时将两个模型的预测结果进行融合」就是干这个的。融合前确认两个模型的输入预处理完全一致否则融合反而更糟。4.5 现象显存爆了训练跑不起来原因512x512 的 patch 加上自注意力和非局部块显存占用比普通 CNN 高不少。解决先把 batch_size 降到 2 甚至 1配合梯度累积模拟大 batch。还不行就把 patch 降到 384 或 256但注意测试时的分块尺寸要同步改训练和测试尺寸必须一致这是文档反复强调的点。5. 进阶玩法把擦除模型导出 ONNX 并做批量验证跑通推理只是第一步真要落地到批量处理试卷得解决两件事一是推理速度二是效果验证。资源包里给了convert_onnx.py这就是提速的入口。导出 ONNX 的典型调用python convert_onnx.py \ --checkpoint ./ckpt/best.pth \ --output ./ckpt/dehw.onnx \ --input_size 512 512 \ --opset 11逻辑说明脚本加载 PyTorch 权重用一张 dummy 输入走一遍前向把计算图固化成 ONNX。导出后可以用 onnxruntime 推理摆脱 PyTorch 依赖部署到没有 GPU 的机器上也能跑。参数说明--input_size必须和训练/测试的 patch 尺寸一致填 512 512。--opset选 11 是兼容性较好的版本太新的 opset 有些推理引擎不认。导出后务必用同一张图分别跑 PyTorch 和 ONNX对比输出差异误差在 1e-3 量级以内才算导出成功。批量验证我一般这么组织把测试集按 512 分块逐块推理后拼回再用 PSNR 和 SSIM 两个指标量化。资源包里PSNRLoss.py已经实现了 PSNR可以直接复用它的计算逻辑别自己重写一套导致口径不一致。验证指标关注点合格参考PSNR整体像素重建质量越高越好横向对比不同权重SSIM结构相似度看印刷体是否完整接近 1 说明结构保留好目视抽查擦除区是否自然、有无残影每批抽 10 张人工过一遍有个血泪经验ONNX 导出后如果发现输出和 PyTorch 对不上八成是某个自定义算子比如非局部块里的 reshape在导出时被简化错了。这时候要么换 opset 版本要么把该模块拆出来单独导出验证。从那以后我每次导出 ONNX都强制走一遍「PyTorch 输出 vs ONNX 输出」的逐像素对比确认误差达标才敢拿去批量跑。这套流程看着麻烦但能省下大量「批量跑完才发现全错」的后悔药。希望帮到你。本文还有配套的精品资源点击获取
返回列表