简介:这份资源面向深度学习与图像处理方向的学习者和开发者,聚焦试卷场景下的手写文字擦除任务,提供从训练到测试的完整工程实现。压缩包共30个文件,以22个Python脚本为核心,涵盖数据加载、损失函数、网络模型与推理逻辑,另含3个Shell脚本、2份readme及说明文档,整体约94KB,结构紧凑便于快速上手。训练采用横向翻转与小角度旋转增强,随机裁剪512×512块,分两阶段优化:先以dice_loss加l1 loss,再仅保留l1 loss。测试环节引入分块与交错分块策略,配合镜像padding和横向镜像增强,并融合两个模型预测结果以提升边缘区域效果。代码按data、loss、model等模块组织,提供compute_mask、train、test等脚本,按说明指定数据与模型路径即可运行。目前已有1397人学习,适合希望复现手写擦除方案、研究分块推理技巧的读者参考。
1. 试卷手写擦除到底在做什么:从一张答题卡说起
改过卷子的老师都有个共识:扫描出来的答题卡上,学生的手写笔迹和印刷题干叠在一起,想拿它做自动批改、题干还原或者电子化归档,第一步就得把那些手写痕迹抹掉,同时不能把印刷的题目、表格线、页码一起擦花。这就是「试卷手写文字擦除」要解决的核心问题——它不是简单的图像去噪,也不是 OCR 识别,而是一个图像到图像的像素级重建任务:输入一张带手写笔迹的试卷图,输出一张只剩印刷内容的干净图。
这个方向属于深度学习里 image inpainting(图像修复)和 document image enhancement(文档图像增强)的交叉地带。和通用去水印、去马赛克不同,试卷场景有几个硬约束:手写笔迹颜色深浅不一、笔画粗细变化大、经常压在印刷字上;印刷内容本身有大量细线条(表格框、下划线、数学符号),一旦擦过头就废了。所以模型不能只学「哪里有笔迹」,还得学「笔迹下面原本是什么」。
这套「源码+模型文件+说明文档」的组合,适合三类人:一是做教育信息化产品、需要批量清洗试卷图像的工程师;二是想拿一个真实文档场景练手图像修复的研究生或转行者;三是手里已经有一批扫描件、想先跑通再决定要不要自训练的团队。下面我按「先搞懂原理和选型 → 再动手跑通 → 最后踩坑和调优」的顺序,把这条链路讲透。
2. 手写擦除的技术选型:为什么多数方案绕不开 GAN 和 U-Net
2.1 从任务本质推导网络结构
手写擦除的输入输出尺寸一致,属于 dense prediction(稠密预测)任务,所以编码器-解码器结构是天然选择。U-Net 的 skip connection 能把浅层的高频细节(印刷字的边缘、表格线)直接传到解码端,避免深层卷积把细线条糊掉——这一点在试卷场景里比在自然图像修复里更关键,因为自然图像丢一点纹理看不出来,试卷丢一根表格线就是事故。
但只用 U-Net 做 L1/L2 回归,输出会偏模糊,手写擦除后经常留下一团灰影。所以主流做法是加一个判别器,用 GAN 的对抗损失逼生成器输出更锐利、更接近真实干净试卷的结果。常见组合是:
- 生成器:U-Net 或带门控卷积(gated convolution)的变体,门控卷积能让网络自己学哪些位置是有效像素、哪些是待修复区域,对不规则笔迹更友好。
- 判别器:PatchGAN,只判断局部 patch 真假,适合文档这种纹理重复度高的图。
- 损失:L1 重建损失 + 对抗损失 + 可选的感知损失(perceptual loss)。
我一般会先跑纯 U-Net + L1 的 baseline,确认数据管线没问题,再加判别器。直接上完整 GAN 容易在早期就震荡,调起来费时间。
2.2 数据从哪来:合成配对样本是主力
真实场景里很难拿到「同一张试卷、有手写版和干净版」的配对数据。所以业内通用做法是合成:拿干净的印刷试卷图,用程序随机叠加手写笔迹,生成 (带笔迹图, 干净图) 配对。手写笔迹来源可以是公开的手写数据集,也可以自己采集。
合成时几个参数直接决定模型上限:
| 参数 | 建议范围 | 说明 |
|---|---|---|
| 笔迹透明度 | 0.6 ~ 1.0 | 太淡模型学不到,太实会盖死印刷字 |
| 笔迹颜色 | 黑/蓝/红随机 | 覆盖圆珠笔、钢笔、红笔批改 |
| 旋转/缩放 | ±15°、0.8~1.2 | 模拟不同书写角度和字号 |
| 叠加位置 | 随机 + 部分强制压字 | 必须有一定比例压在印刷内容上 |
| 模糊/噪声 | 轻度高斯 | 模拟扫描件质量差异 |
提示:合成数据里「笔迹压在印刷字上」的样本比例不能太低,否则模型遇到真实压字场景会直接翻车。我一般让这个比例占到 40% 以上。
2.3 源码和模型文件到手后先看什么
拿到一个「源码+模型文件+说明文档」的包,别急着跑训练。先按这个顺序确认:
- 模型文件格式:是
.pth、.pt还是.onnx。PyTorch 权重需要对应的网络定义代码,ONNX 可以直接推理。 - 输入输出规格:看说明文档里的输入尺寸(常见 256×256 或 512×512)和归一化方式。
- 依赖版本:重点看 PyTorch / torchvision 版本,版本不匹配加载权重会报 key 错误。
- 有没有推理脚本:优先找
inference.py或test.py,先跑单张图验证,再碰训练。
这一步能省掉后面大量「以为是模型问题、其实是环境问题」的排查时间。
3. 本地跑通推理:从环境配置到单张试卷出图
3.1 环境配置的最小命令集
先建独立环境,避免和系统里的其他深度学习项目打架。下面以 conda 为例:
# 创建环境,python 版本按说明文档要求,常见 3.8~3.10 conda create -n paper_erase python=3.10 -y conda activate paper_erase # 安装 PyTorch,CUDA 版本按自己显卡驱动选,这里以 cu118 为例 pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 # 安装图像处理和推理常用依赖 pip install opencv-python pillow numpy tqdm逻辑说明:PyTorch 单独用官方 index 装,避免 pip 默认源拉到 CPU 版。opencv-python用于读写和预处理图像,tqdm用于批量推理时看进度。参数上,CUDA 版本要和nvidia-smi显示的驱动兼容,驱动太旧就降 CUDA 版本,别硬上。
3.2 加载模型并跑单张图
假设源码里网络定义在models/unet.py,权重是weights/eraser.pth。写一个最小推理脚本:
import torch import cv2 import numpy as np from models.unet import UNet # 按实际路径改 # 1. 设备选择:有卡用卡,没卡用 CPU(慢但能跑) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") # 2. 实例化网络,结构必须和训练时一致,否则权重加载会报错 model = UNet(in_channels=3, out_channels=3).to(device) # 3. 加载权重,map_location 保证在 CPU 上也能加载 GPU 存的权重 state = torch.load("weights/eraser.pth", map_location=device) model.load_state_dict(state) model.eval() # 推理模式,关掉 dropout 和 batchnorm 更新 # 4. 读图并预处理:缩放到模型输入尺寸,转 tensor,归一化到 [-1,1] img = cv2.imread("test_paper.jpg") img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img = cv2.resize(img, (512, 512)) tensor = torch.from_numpy(img).float().permute(2, 0, 1) / 127.5 - 1.0 tensor = tensor.unsqueeze(0).to(device) # 加 batch 维度 # 5. 前向推理,不计算梯度省显存 with torch.no_grad(): output = model(tensor) # 6. 后处理:反归一化回 [0,255],转回 numpy 存图 output = (output.squeeze(0).permute(1, 2, 0).cpu().numpy() + 1.0) * 127.5 output = np.clip(output, 0, 255).astype(np.uint8) cv2.imwrite("clean_paper.jpg", cv2.cvtColor(output, cv2.COLOR_RGB2BGR))逻辑说明:整个流程是「读图 → 预处理 → 前向 → 后处理 → 存图」。关键参数有三个:输入尺寸必须和训练时一致(这里 512),归一化方式必须和训练时一致(这里 [-1,1]),model.eval()不能漏。如果加载权重报Missing key(s)或Unexpected key(s),八成是网络定义和权重不匹配,先核对in_channels、层数、是否有门控卷积。
3.3 批量推理和结果检查
单张跑通后,改成批量处理整个文件夹:
import os from glob import glob os.makedirs("output", exist_ok=True) for path in glob("input/*.jpg"): img = cv2.imread(path) # ... 复用上面的预处理和推理逻辑 ... name = os.path.basename(path) cv2.imwrite(f"output/{name}", result)跑完别只看一两张就下结论。重点检查三类图:笔迹压在印刷字上的、有表格线的、印刷字本身很细的(比如数学公式)。这三类是最容易暴露模型短板的。如果表格线被擦断,说明模型对细结构保留不足,后面调优要往这个方向使劲。
4. 自己训练时的数据管线和参数设置
4.1 合成数据的脚本骨架
如果包里的模型在你的场景效果不够,就得自己微调或重训。核心是造配对数据:
import cv2 import numpy as np import random def synthesize(clean_img, handwriting_imgs): """把随机手写笔迹叠加到干净试卷上""" h, w = clean_img.shape[:2] canvas = clean_img.copy() # 随机叠加 3~8 个手写片段 for _ in range(random.randint(3, 8)): hw = random.choice(handwriting_imgs) # 随机缩放和旋转 scale = random.uniform(0.8, 1.2) hw = cv2.resize(hw, None, fx=scale, fy=scale) angle = random.uniform(-15, 15) M = cv2.getRotationMatrix2D((hw.shape[1]//2, hw.shape[0]//2), angle, 1) hw = cv2.warpAffine(hw, M, (hw.shape[1], hw.shape[0])) # 随机位置 x = random.randint(0, max(0, w - hw.shape[1])) y = random.randint(0, max(0, h - hw.shape[0])) # 透明度混合 alpha = random.uniform(0.6, 1.0) roi = canvas[y:y+hw.shape[0], x:x+hw.shape[1]] mask = hw < 200 # 手写像素掩码 roi[mask] = (alpha * hw[mask] + (1-alpha) * roi[mask]).astype(np.uint8) return canvas逻辑说明:alpha控制笔迹深浅,mask只让手写像素参与混合,避免把背景也叠上去。参数上,叠加数量、缩放范围、旋转角度都要按真实试卷的书写密度调。合成完记得把 (带笔迹图, 干净图) 成对存好,训练时同步做增强。
4.2 训练参数怎么设
以 U-Net + PatchGAN 为例,常用配置:
| 参数 | 建议值 | 说明 |
|---|---|---|
| 输入尺寸 | 256 或 512 | 512 显存占用约 4 倍,效果通常更好 |
| batch size | 4~8(512)/ 8~16(256) | 按显存调,别爆 |
| 学习率 | 2e-4 | Adam,生成器和判别器可分开设 |
| 优化器 | Adam(β1=0.5) | β1 设 0.5 是 GAN 惯例,比默认 0.9 稳 |
| L1 权重 | 100 | 重建损失权重,太高会模糊,太低会失真 |
| 训练轮数 | 50~200 | 看验证集,别死磕轮数 |
注意:GAN 训练早期判别器容易过强,导致生成器梯度消失。常见做法是判别器学习率设成生成器的一半,或者前几轮只训生成器。
4.3 验证指标和肉眼检查
自动指标用 PSNR 和 SSIM,但这两个在文档场景参考价值有限——PSNR 高的图可能表格线还是断的。我的习惯是每存一次 checkpoint,就固定抽 10 张验证图拼成对比图(原图 / 带笔迹 / 模型输出 / 干净真值),肉眼过一遍。断线、灰影、印刷字被误擦,这三类问题自动指标看不出来,只能靠眼睛。
5. 避坑与排查:手写擦除最容易翻车的五个地方
5.1 输出整片发灰,像蒙了层雾
现象:擦除后手写是没了,但整张图灰蒙蒙,印刷字对比度下降。 原因:L1 损失权重过高,模型倾向于输出「平均色」来降低像素误差,导致模糊。 解决:降低 L1 权重,加入对抗损失或感知损失;也可以在损失里对印刷字区域加权,逼模型保住高对比边缘。
5.2 表格线和下划线被擦断
现象:手写擦干净了,但表格框出现缺口,数学下划线断成几截。 原因:训练数据里细线条样本不足,或者网络下采样太深,浅层细节没传到位。 解决:合成数据时增加细线条场景;检查 U-Net 的 skip connection 是否完整;必要时用膨胀后的笔迹掩码做局部修复,而不是全图重建。
5.3 加载权重报 key 不匹配
现象:load_state_dict抛Missing key(s)或Unexpected key(s)。 原因:网络定义和权重版本不一致,常见于源码更新过但权重没换,或者in_channels对不上。 解决:打印state.keys()和model.state_dict().keys()对比,找出差异层;确认是否有多 GPU 训练留下的module.前缀,有的话用state = {k.replace('module.', ''): v for k, v in state.items()}去掉。
5.4 推理速度慢到没法批量用
现象:单张 512 图要好几秒,几千张跑一天。 原因:没上 GPU、没开半精度、或者模型本身太重。 解决:确认torch.cuda.is_available()为 True;推理时用with torch.autocast('cuda')开混合精度;如果还慢,把模型导出成 ONNX 或 TensorRT,推理能快 2~5 倍。
5.5 换一批试卷就效果暴跌
现象:在 A 学校试卷上很好,换 B 学校的扫描件就一堆残留。 原因:训练数据分布太窄,扫描分辨率、纸张底色、笔迹颜色都和训练集差异大。 解决:拿新场景的图做少量微调(几十张就够),或者做域增强——训练时随机调亮度、对比度、加噪声,提升泛化。
6. 进阶技巧:用掩码引导和分块推理把效果再拉一档
跑通基础流程后,想让效果更稳,有两个我常用的技巧。
第一个是掩码引导推理。纯生成式擦除是「盲擦」,模型不知道哪里该擦。如果能先用手写检测(一个轻量分割网络)生成笔迹掩码,再把掩码作为额外输入通道喂给擦除网络,模型就能把注意力集中在笔迹区域,印刷字被误擦的概率明显下降。实现上就是把 U-Net 的in_channels从 3 改成 4,第四通道是掩码,训练时用合成笔迹的掩码做监督。代价是要多训一个检测模型,但擦除质量提升值得。
第二个是分块推理。512 的模型直接吃 2000 多像素的扫描件会被迫缩放,细线条全糊。做法是把大图切成带重叠的 512 块,逐块推理再拼回去,重叠区域做加权融合。代码骨架:
def tiled_inference(model, img, tile=512, overlap=64): h, w = img.shape[:2] output = np.zeros_like(img, dtype=np.float32) weight = np.zeros((h, w, 1), dtype=np.float32) step = tile - overlap for y in range(0, h, step): for x in range(0, w, step): y2, x2 = min(y+tile, h), min(x+tile, w) y1, x1 = max(0, y2-tile), max(0, x2-tile) patch = img[y1:y2, x1:x2] # ... 预处理 + 模型推理 ... output[y1:y2, x1:x2] += result weight[y1:y2, x1:x2] += 1 return (output / np.maximum(weight, 1)).astype(np.uint8)逻辑说明:overlap是重叠像素,设太小拼接处会有缝,设太大浪费算力,一般取 tile 的 1/8 到 1/4。weight记录每个像素被覆盖次数,最后取平均,边缘过渡自然。参数上,tile 要和模型训练尺寸一致,overlap 至少 32。
验证效果别只看 PSNR。我的习惯是固定三张「地狱难度」图——笔迹压公式、红笔批改压表格、浅色铅笔字——每次改动都拿它们对比。这三张过了,基本就稳了。
最后说个血泪教训:别一上来就追求端到端大模型。我早期直接上 1024 分辨率的大网络,显存爆了不说,训了两天效果还不如 512 的小模型加掩码引导。先把数据管线和评估流程做扎实,模型大小是最后才调的事。希望帮到你。
本文还有配套的精品资源,点击获取