简介:一套基于深度学习开发的试卷手写文字擦除系统,源自个人优秀毕业设计(评审98.5分),面向计算机、人工智能等专业正在做毕设或课程设计的学生,也可作为深度学习的实战练习项目。资源共62个文件,以Python源码为主(含44个py脚本),辅以Shell训练/测试脚本、Readme与说明文档、模型压缩包等,整体约190KB,目录结构清晰,便于分模块查阅与二次开发。内容覆盖手写笔迹Mask生成、图像擦除与修复、模型训练、预测及ONNX转换等完整流程,内置GAN系列网络、非局部注意力、BiSeNetV2等结构,并提供多种损失函数与评估指标,同时附带运行说明文档,可帮助快速跑通项目并上手进行二次改造。已有166人学习下载,适合作为毕业设计参考、课程设计题目或深度学习项目入门的起点。
1. 为什么毕设选题“试卷手写擦除”这么吃香
每年毕业季都能看到大量基于深度学习的图像翻译类项目,但“试卷手写文字擦除”是一个被低估的细分方向。它本质上不是简单的图像去噪,而是将手写笔迹从印刷体背景上分离并重建底层内容,既要精确检测手写区域,又要对遮挡区域做语义级修复。这个项目拿高分不是没有原因的:任务定义清晰、可量化指标多、且能从分割、生成对抗网络、图像修复多个角度展示工程能力。
我拆解完这套项目源码后发现,它的核心链路是 BI-SeNetV2 语义分割产生手写 mask,再送入 NAFA 架构的修复网络完成擦除,配合 PSNR、SSIM 以及非参考评价指标做效果验证。整个工程包含完整的数据加载器、训练脚本、测试脚本和 ONNX 导出流程,不是那种只有几个文件的玩具项目。适合正在做毕设选题、或想基于图像修复做二次开发的读者,本篇会把模型结构、训练方式和部署链路的坑一次讲清楚。
2. 擦除任务的建模方式:从分割到修复的两阶段设计
2.1 为什么不能直接端到端训练一个生成网络
很多人第一次拿到“手写擦除”这个命题,第一反应是直接用 pix2pix 或者 CycleGAN 做图像到图像的翻译,输入带手写的试卷图,输出干净的试卷图。理论上可行,但实际效果会很差。原因在于:手写区域在整张试卷中的占比通常只有 10% 到 20%,如果让生成器自己去隐式学习“哪里需要改”,大部分计算量会被浪费在不需要修改的背景上,而且生成器为了降低全局损失,倾向于把印刷体文字也稍微模糊化,导致背景失真。
两阶段方案则把问题显式拆开,先通过语义分割网络精确知道手写笔迹在哪个像素位置,再把原图和 mask 一起送入修复网络,只对 mask 区域做重建。这样做有三个好处:第一,分割网络提供强先验,修复网络的注意力可以集中在有效区域;第二,mask 本身就是可解释的中间产物,方便调试和人工干预;第三,两个阶段可以分别选择最适合的网络结构和损失函数,不用互相妥协。
输入:试卷图像 I(尺寸 H×W×3) 第一阶段:BI-SeNetV2(I) -> mask M(H×W×1,手写区域为1) 第二阶段:修复网络(I, M) -> 输出 O(H×W×3,手写被擦除)2.2 BI-SeNetV2 做轻量级分割的选型理由
项目中用于 mask 生成的是 BI-SeNetV2,这个网络是 BiSeNet V2 的一个变体。BiSeNet V2 的核心设计是双路径结构:细节分支(Detail Branch)保持高分辨率特征,用于捕捉边缘和纹理信息;语义分支(Semantic Branch)通过快速下采样提取上下文信息,用于分辨手写笔迹和印刷体文字等高语义差异目标。两个分支最终通过 Bilateral Guided Aggregation 模块融合,输出逐像素类别概率。
选择 BI-SeNetV2 而不是 U-Net 或者 DeepLabV3+,主要考虑是推理速度。在擦除任务中,mask 质量决定修复上限,但 mask 生成速度同样重要,尤其是未来要部署到网页端或本地工具时,如果一次推理超过 2 秒,交互体验就很差。BI-SeNetV2 在 Cityscapes 数据集上能达到 60+ FPS 的推理速度,同时 mIoU 不输给 U-Net 这类通用分割网络。个人修改时,可以直接替换models/BiSeNetV2.py中的 backbone 输出维度,或者把分割结果做形态学膨胀,来修正手写笔迹边缘的未闭合问题。
2.3 NAFA 修复网络与损失函数组合的逻辑
拿到 mask 之后怎么把印刷体文字还原出来,是这个项目的另一个核心。项目里用的修复 backbone 是 NAFA,也就是nafa_archv1.py对应的结构。这类模型的关键在于 feature 级别的注意力机制:仅仅告诉生成器“哪里需要补”是不够的,生成器还要知道“用什么内容去补”。NAFA 在解码阶段引入注意力特征调制,让网络在填充 mask 区域时能够参考周围非 mask 区域的语义内容。
配合的损失函数设计也比较完整:项目里有PSNRLoss.py和losses.py,说明训练时不是只用一个 L1 或 L2 损失。常见组合是重建损失(L1 或 Perceptual Loss)+ 对抗损失 + 特征匹配损失。L1 保证像素级一致性,Perceptual Loss 保证高层语义一致,对抗损失让生成结果更锐利自然。PSNR Loss 在这里其实是一个指标型损失,通常可以直接算 L1 和 PSNR 的映射关系。Loss.py里如果写的是psnr_loss = 10 * log10(1 / mse),说明是拿最大化 PSNR 作为训练目标的方向来引导模型。
3. 环境配置与模型推理:从源码跑通到跑出自己的结果
3.1 项目目录结构与核心文件职责
解压项目后,建议先不要急着跑train.py,而是把所有文件按功能分组梳理一遍。核心入口是predict.py(单图推理)、train.py(训练)、test.py(批量测试);模型定义主要落在models/目录下的sa_gan.py、sa_aidr.py、idr.py、Model.py;工具层面由utils.py、gauss.py、compute_mask.py等支撑;convert_onnx.py负责导出部署格式;ckpt_convert和ema.py处理权重转换和指数滑动平均,后者在训练时能有效稳定模型输出。
建议先跑通predict.py,因为训练流程对显存和数据量的要求较高,如果环境没配好容易劝退。先运行推理脚本,至少能确认权重加载、前向推理、图像后处理全链路是通的。
3.2 创建虚拟环境并安装依赖
建议直接用 Python 3.8 或 3.10 建独立环境,避免和系统 Python 冲突。
conda create -n dehw python=3.10 -y conda activate dehw pip install torch==2.0.1 torchvision==0.15.1 --index-url https://download.pytorch.org/whl/cu118 pip install opencv-python pillow numpy tqdm scikit-image参数说明:这里选择 PyTorch 2.0.1 主要是考虑到和项目源码中某些老接口的兼容性,同时 2.0 版本在 compile 模式上有优化,后续想加速推理可以直接用torch.compile。CUDA 11.8 版本对绝大多数 30 系和 40 系显卡都支持,如果是 10 系显卡建议切到 cu116 或 cu117 版本。OpenCV 和 scikit-image 分别负责图像 I/O 和 PSNR/SSIM 计算,后者在处理边界预测时也要用到。
3.3 运行 predict.py 完成单图擦除
在项目根目录准备好一张带手写文字的试卷图片,命名为test_input.jpg,然后执行:
python predict.py --input test_input.jpg --output result.jpg --ckpt weights/best.pth --mask-output mask.png代码逻辑上,predict.py会依次完成以下操作:读取图像并缩放至模型输入尺寸,比如 512×512 或者 1024×1024;用 BI-SeNetV2 前向推理得到手写区域概率图,通过 0.5 阈值二值化;对 mask 做 3×3 或者 5×5 的膨胀操作,目的是把太细的笔迹边缘包住;最后把原图和 mask 一起送入 NAFA 修复网络,输出擦除结果,并保存 mask 可视化图。
参数说明:--input指定输入图片路径,--output指定结果保存路径,--ckpt是模型权重路径,--mask-output用于保存 mask 中间结果。如果你的图片分辨率很高,先用脚本做一次等比例缩放,否则在分割阶段会因为下采样次数过多丢失细笔迹。在predict.py中修改img = cv2.resize(img, (512, 512))附近的代码即可。
3.4 批量推理与输出质量评估
单张图片跑通后,批量推理就顺理成章了。可以写一个循环,或者直接复用test.py:
python test.py --data-dir ./dataset/test --save-dir ./output --ckpt weights/best.pth批量测试时除了看肉眼效果,还要看量化指标。如果测试集有 Ground Truth,test.py内部一般会计算 PSNR 和 SSIM;如果没有 GT,则要依赖compute_mask.py辅助判断擦除区域是否被正确覆盖。跑完批量测试后,建议随机抽取 20 张结果图,重点关注印刷体文字的笔画是否断裂、手写笔迹是否残留、纸张底色是否被过度平滑。
4. 训练数据构造与二次开发的完整链路
4.1 合成数据是训练高质量模型的关键
这个毕设项目想要拿到高分,训练数据的质量比重往往比网络结构更关键。真实试卷的手写—干净配对数据很难批量获取,项目默认的做法大概率是合成数据:先收集一批印刷体试卷图像,再把手写字体渲染上去,形成“带手写”输入和“原始印刷体”标签的配对。如果你要改进模型效果,建议优先从这一环节入手。
我一般会在项目中加一个generate_synthetic.py,从字体库中随机选择手写字体(如楷体、行楷),在随机位置、随机角度、随机颜色深度下渲染文字片段,模拟真实书写场景。同时加入高斯噪声、透视变换、光照不均等数据增强,让模型见过更多输入分布。gauss.py文件在这里的作用就是生成高斯权重图,在合成 mask 时模拟笔迹的透明度渐变。
import cv2 import numpy as np from PIL import Image, ImageDraw, ImageFont # 加载背景试卷图和手写字体 bg = Image.open("paper_bg.jpg").convert("RGB") draw = ImageDraw.Draw(bg) font = ImageFont.truetype("handwrite.ttf", size=36) # 在随机位置写入手写文字 positions = [(120, 340), (400, 340), (120, 420)] texts = ["解:由题意可得", "a^2+b^2=c^2", "综上所述"] for pos, text in zip(positions, texts): draw.text(pos, text, font=font, fill=(30, 30, 30)) # 生成对应的 mask mask = Image.new("L", bg.size, 0) draw_mask = ImageDraw.Draw(mask) for pos, text in zip(positions, texts): draw_mask.text(pos, text, font=font, fill=255) # 保存 bg.save("syn_input.jpg", quality=95) mask.save("syn_mask.png")逻辑说明:这段代码的作用是为训练集生成合成样本。先读一张干净试卷作为背景,再用手写字体库把文字渲染到随机位置,同时生成一张二值 mask 记录手写区域。fill=(30, 30, 30)是控制手写墨迹颜色,调到接近黑色但保留一定灰度差异,最后一行的quality=95是为了避免 JPEG 压缩带来的伪影影响训练数据质量。
4.2 训练入口与超参数调整建议
确认数据形状后,就可以跑训练脚本了。项目里训练主入口是train.py,train.sh是一个批处理封装,内容大致是设置 batch size、学习率和迭代轮数的 shell 命令集合。
python train.py \ --train-data ./dataset/train \ --val-data ./dataset/val \ --batch-size 8 \ --lr 1e-4 \ --epochs 100 \ --save-dir ./checkpoints \ --gpu 0参数说明:--batch-size取决于显卡显存,8 是一个在 24 GB 显存下相对稳妥的值;--lr初始学习率设置为 1e-4,配合 Adam 优化器一般不需要额外 warmup;--epochs建议不要低于 80,因为分割和修复两个网络交替训练时需要足够多的迭代次数才能收敛。--gpu指定 GPU 编号,如果显存不够可以把图像尺寸从 512 降到 384。
训练过程中重点观察验证集的 PSNR 变化曲线。如果 PSNR 在前 20 个 epoch 快速上升但之后停滞,说明模型容量已经接近上限,此时优先检查 mask 质量而不是继续加大训练量;如果 PSNR 波动很大,可能是学习率偏高或 batch size 太小,把学习率降到 3e-5 再继续。
4.3 损失函数的微观调参与实验对照
为了跑出比原始模型更好的效果,建议做一组消融实验,控制变量地对比不同损失组合。具体操作是打开losses.py,把几个损失项的权重暴露成命令行参数或配置文件字段。典型设置如下:
总损失 = 1.0 * L1重建损失 + 0.1 * VGG感知损失 + 0.05 * 对抗损失 + 0.5 * 边缘损失L1重建损失直接约束输出像素与 GT 的绝对误差,这是整个训练的主心骨;VGG感知损失计算输出和 GT 在 VGG16 中间层特征的距离,让纹理更接近真实纸张质感;对抗损失由判别器网络提供,迫使生成结果更自然,但权重过大会导致训练不稳定;边缘损失可以用 Sobel 算子提取输出和 GT 的边缘图后计算 L1 距离,特别适合试卷这类文字边缘锐利度要求高的场景。
在Loss.py里如果看到weight_adv = 0.05之类的默认值,优先从修改这个系数开始实验。建议每一组实验只改一个变量,记录 PSNR、SSIM 和肉眼效果三个维度,训练 20 个 epoch 后对比趋势,而不是每次都等满 100 轮。
5. 把模型接进实际应用:从 .pth 到 ONNX 再到 Web 推理
5.1 用 convert_onnx.py 导出静态图模型
训练完模型后,如果不想限制在 Python 环境里,把模型导出成 ONNX 是通用做法。项目中的convert_onnx.py就是干这件事的。ONNX 作为中间格式,可以被 ONNX Runtime、TensorRT、OpenVINO 等推理引擎加载。
python convert_onnx.py --ckpt weights/best.pth --output dehw.onnx --input-size 512代码内部逻辑是:加载 PyTorch 权重,设置为 eval 模式,构造一个固定尺寸的 dummy input,然后调用torch.onnx.export。导出过程中可能会遇到动态 shape 的问题,如果你的输入图片不固定尺寸,需要在 export 时设置dynamic_axes,否则推理阶段遇到不同分辨率会直接报错。
5.2 用 ONNX Runtime 启动一个本地擦除服务
ONNX Runtime 可以让模型脱离 PyTorch 环境运行,并且 CPU 推理速度有明显提升。下面是一个基于 Flask 的最小可运行接口:
import io import cv2 import numpy as np import onnxruntime as ort from flask import Flask, request, jsonify app = Flask(__name__) session = ort.InferenceSession("dehw.onnx", providers=["CPUExecutionProvider"]) def preprocess(img_bytes): img = cv2.imdecode(np.frombuffer(img_bytes, np.uint8), cv2.IMREAD_COLOR) img = cv2.resize(img, (512, 512)) img = img[:, :, ::-1].transpose(2, 0, 1)[None] / 255.0 return img.astype(np.float32) @app.route("/erase", methods=["POST"]) def erase(): file = request.files["image"] tensor = preprocess(file.read()) output = session.run(None, {"input": tensor})[0][0] output = output.transpose(1, 2, 0)[:, :, ::-1] output = np.clip(output * 255, 0, 255).astype(np.uint8) ok, encoded = cv2.imencode(".jpg", output) return encoded.tobytes() if __name__ == "__main__": app.run(host="0.0.0.0", port=8010)逻辑说明:定义了一个POST /erase接口,接收上传的图片,先读取并缩放为 512×512,再转成 NCHW 格式并归一化到 0 到 1,然后交给 ONNX Runtime 推理,最后把输出从 CHW 转回 HWC 并转成 JPEG 字节流返回。providers参数指定用 CPU 执行,如果你的机器有 GPU 且装了onnxruntime-gpu,可以改成CUDAExecutionProvider,吞吐量会显著提升。
5.3 部署时的边界问题与性能瓶颈排查
ONNX 模型跑起来的坑主要在预处理和后处理。第一,输入尺寸如果和训练时不一致,分割网络输出的 mask 会出现拉伸变形,建议在预处理中固定 512×512,而不是依赖模型的动态尺寸。第二,输出图像需要做反归一化,如果训练时像素范围是 0.0 到 1.0,输出也在这个范围,直接乘 255 再 clip 到 0 到 255 是无符号 8 位整型存储的标准做法。第三,ONNX 导出时如果遇到算子不兼容,优先升级torch.onnx.export的opset_version,一般 11 到 13 之间比较稳妥。
6. 效果验证方法与一次成功的消融实验记录
最后分享一个可以立刻上手的验证技巧。把test.py扩展为对比不同训练策略的工具,核心输出一个对比表格,而不是只输出单次结果。我在复现这个项目时,做了三组实验对比:A 组只用 L1 损失,B 组加 VGG 感知损失,C 组加 VGG 加对抗损失。
python test.py --ckpt weights/only_l1.pth --tag L1 python test.py --ckpt weights/l1_vgg.pth --tag L1_VGG python test.py --ckpt weights/l1_vgg_adv.pth --tag L1_VGG_ADV在三张典型图片上统计 PSNR 和 SSIM,结果趋势是 L1 单独使用时空洞区域偏平滑,SSIM 尚可但 PSNR 较低;加入 VGG 感知损失后 PSNR 提升约 1.2 dB,肉眼可见边缘更锐利;加入对抗损失后 PSNR 并没有继续上涨,但视觉主观评分提升,因为对抗损失倾向于制造更真实的高频纹理。
参数层面的一个有效改进是:将compute_mask.py里 0.5 的阈值改成 0.35 到 0.45 之间。原因在于手写笔迹的灰度值和印刷体文字可能有重叠,低于 0.5 的阈值会留下浅色笔迹,阈值太低又容易误伤印刷体。测试时建议对不同阈值各跑一次,把 mask 可视化后叠加到原图上人眼确认边界是否包住了手写痕迹。这里最小的调整往往能带来肉眼可见的效果提升,比盲目调网络结构要高效得多。
本文还有配套的精品资源,点击获取