☰
试卷手写擦除实战:U-Net与GAN源码模型推理及训练调优指南
2026/9/28 16:39:00 网站建设 项目流程

简介:这份资源面向深度学习与图像处理方向的学习者和开发者,聚焦试卷场景下的手写文字擦除任务,提供从训练到测试的完整工程实现。压缩包共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 源码和模型文件到手后先看什么

拿到一个「源码+模型文件+说明文档」的包,别急着跑训练。先按这个顺序确认:

  1. 模型文件格式:是.pth、.pt还是.onnx。PyTorch 权重需要对应的网络定义代码,ONNX 可以直接推理。
  2. 输入输出规格:看说明文档里的输入尺寸(常见 256×256 或 512×512)和归一化方式。
  3. 依赖版本:重点看 PyTorch / torchvision 版本,版本不匹配加载权重会报 key 错误。
  4. 有没有推理脚本:优先找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 或 512512 显存占用约 4 倍,效果通常更好
batch size4~8(512)/ 8~16(256)按显存调,别爆
学习率2e-4Adam,生成器和判别器可分开设
优化器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 的小模型加掩码引导。先把数据管线和评估流程做扎实,模型大小是最后才调的事。希望帮到你。

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

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

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

立即咨询