☰
基于深度学习的试卷手写擦除:两阶段训练与分块推理实战
2026/9/28 23:35:31 网站建设 项目流程

简介:这份资源面向深度学习与图像处理方向的学习者和开发者,提供一套基于深度学习的试卷手写文字擦除完整实现方案,可用于试卷还原、文档图像净化等场景,适合具备一定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数据加载,负责读图、增强、切 patch
loss/(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 python=3.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_size(512 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 输出」的逐像素对比,确认误差达标才敢拿去批量跑。这套流程看着麻烦,但能省下大量「批量跑完才发现全错」的后悔药。希望帮到你。

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

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

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

立即咨询