手写文字擦除实战:从mask生成到深度学习模型训练与部署
2026/9/13 16:24:58 网站建设 项目流程

简介:面向大学生竞赛与深度学习开发者,这是一套手写文字擦除赛题冠军方案,基于Python与深度学习技术,针对试卷中红蓝黑多色手写字、手绘线段与污渍、手写印刷字重叠等复杂场景,实现文字擦除与背景修复。压缩包共30个文件,整体约95KB,以22个Python源码为主,覆盖数据加载、网络结构、损失函数与PSNR指标计算、训练推理及模型转换等环节;同时包含3个shell脚本、2份readme说明及txt/md文档,按流程即可复现和二次开发。已有337人学习下载。方案额外提供了数据集划分细节与官方赛题特征总结,结合模型文件和说明文档,可帮助读者快速跑通完整流程、理解榜首技术思路,并迁移至答题卡清洁、文档图像修复等任务。

1. 手写文字擦除到底解的是哪一类问题

手写文字擦除,业界通常叫 Handwritten Text Removal,和通用图像修复(inpainting)不同,它不要求自由创作补全,而是要在一张扫描稿上精准去掉手写笔迹,同时保留印刷体、表格线和纸张纹理。它广泛用于档案电子化、试卷去笔迹、表单重建和 OCR 数据清洗。传统形态学和阈值方法处理不了半透明铅笔字和交叉笔画,深度学习方案是当前性价比最高的路径。标题里的下载即用 python 源码,实质是把数据集、mask、训练、推理打包成一条可复现流程。对工程人员来说,核心是先确认数据约定与模型 io,再谈调参和换骨干。

2. 数据组织与手写笔迹 mask 的生成方式

2.1 确定数据集的目录约定与标注格式

下载即用的方案包里,数据部分一般不像竞赛作业那样只给 JPG,而是按「训练/验证/测试」和「原图/掩码/干净图」分成三个平行目录。拿到一个包,先执行 tree 命令看结构,而不是急着配环境。

tree project_root -L 2

最常见的目录结构是三件套并行:

目录内容典型格式
train/images带手写的扫描稿.png / .jpg
train/masks手写区域二值图.png(0 或 255)
train/clean对应的干净底图.png / .jpg
val / test同为三件套同上

有的包会额外提供标注 JSON 或 XML,对应的是「检测式」标注方式,坐标是多边形或矩形框,需要离线转换成 mask;有的直接给像素级二值图,是给修复模型用的。无论哪种,最终都要归一成 0/255 的 uint8 图,PyTorch 的 DataLoader 只需要它和 image、target 对齐即可参与计算。

mask 的语义值得花时间确认:255 通常表示手写区域,0 表示背景,但也有反过来的包。很多训练事故就出在颜色反转上,模型把背景全擦了。收到数据后先np.unique(mask)看一下数值分布;如果出现中间值,说明 mask 经过抗锯齿或 JPEG 压缩,需要重新阈值化。

2.2 从多边形标注生成像素级 mask

有些方案包不带预处理好的 mask,只给 JSON/XML 标注,需要在训练前离线生成。用 OpenCV 的 fillPoly 几分钟就能做完,关键是把坐标系的宽高对齐到原图尺寸,别用反。

import cv2 import numpy as np import xml.etree.ElementTree as ET def xml_to_mask(xml_path, img_w, img_h): tree = ET.parse(xml_path) root = tree.getroot() mask = np.zeros((img_h, img_w), dtype=np.uint8) for obj in root.iter('object'): pts = [] for pt in obj.iter('pt'): x = int(round(float(pt.find('x').text))) y = int(round(float(pt.find('y').text))) pts.append([x, y]) if len(pts) >= 3: cv2.fillPoly(mask, [np.array(pts, dtype=np.int32)], 255) # 防止标注越界导致后续训练读取异常 return cv2.copyMakeBorder(mask, 0, 0, 0, 0, cv2.BORDER_CONSTANT)

这段代码做的事很直接:把 XML 里每个 object 下的多边形顶点解析出来,转成 np.int32 点数组,交给 fillPoly 填充;填充值 255 是约定。解析时要注意 XML 里 x、y 可能是字符串或带小数点,float() 再 round 才安全。如果标注出现坐标越界(扫描仪裁剪导致),fillPoly 会画出错误区域,需要先 clip 到[0, w-1][0, h-1]

如果标注是矩形框而不是多边形,常见做法是把四个角连成矩形填充,然后对 mask 做一次 dilate,膨胀 3~5 像素,目的是罩住真实笔迹的浅色毛边。手写笔迹的灰度是渐变的,标注框往往只包含深色核心,不膨胀会导致擦除后残留一圈淡色。

2.3 数据增强的 3 个关键操作

手写擦除任务的增广不能只做随机翻转和裁剪,还要考虑纸面的真实变化:

  1. 弹性形变(Elastic Deformation):模拟纸张皱褶和局部弯曲。用 OpenCV 的 remap 配合一个随机位移场即可。注意 image、mask、clean 三张图必须用同一个位移场,否则 mask 和内容对不上。
  2. 亮度与色温扰动:不同扫描仪、不同光照下纸面色泽差异很大。训练时以 50% 概率给图像乘一个 0.85~1.15 的随机因子,再叠加一个较小的颜色抖动。
  3. 随机透视:模拟手机翻拍而非扫描仪平扫。用 cv2.getPerspectiveTransform 生成轻微透视变换,同样作用在三件套上。

弹性形变最稳妥的实现是调用 imgaug 里的 ElasticTransformation,或者自己维护一个 warp 函数。如果不想引入额外依赖,用 random_crop + rotation 代替也行,但效果会差一些,因为真实扫描件经常有局部弯曲,简单的全局变换模拟不了。

2.4 踩坑:背景纹理被当手写擦掉

第一个坑是印章和印刷体被当成手写。注意确认数据标注里手写和印章是否分开;有的数据集里印章也算标注对象,但语义是「要去除的东西只有笔迹」。

第二个坑是 mask 膨胀过度。膨胀多了会把印刷体边缘也罩进 mask,模型被迫重画印刷体,产生字形畸变。一个可行的检查办法是:从训练集中随机抽几张图,把 mask 以红色叠加显示在原图上,肉眼确认 mask 边界和笔迹边缘的贴合程度。如果发现 mask 明显盖过印刷体,把膨胀核从 5x5 降到 3x3 或去掉。

第三个坑是训练 crop 比例。修复类模型对大空洞很敏感:如果一张 crop 里有 40% 以上是 mask,梯度会被无意义的填充主导,模型学不到「保持外部不变」的基本任务。一般做法是在 Dataset 里控制 crop 区域 mask 占比在 10%~30%,超了就重新采样。

提示:拿到包后,先跑数据可视化脚本看 3 组样本(原图、mask、clean),确认 mask 方向和值域,再进训练。这一步能省掉后续大部分疑难杂症。

3. 模型结构与骨干选型:修复式还是分割式

3.1 修复式(inpainting)路线的网络骨架

手写擦除最常见的实现是把任务建模为 mask-conditional inpainting:输入是「原图 + 二值 mask」,输出是重建后的干净图。骨干网络的选择顺序,我会按下面的规律去试:

  • 最保守的方案:U-Net + 普通卷积。因为任务输入输出分辨率一致,U-Net 的跳跃连接能保留低频背景信息,在 mask 较小的样本上效果稳定。
  • 更符合任务特性的方案:部分卷积或门控卷积。普通卷积在 mask 边界会把 mask 内外的特征混在一起,导致边界发灰;partial convolution 每层只在有效像素上做卷积,同时学习一个 mask 更新规则,把「该不该在这个位置填充」建模成可学习信号。
  • 追求高真实感的方案:CoModGAN 风格的生成器 + 判别器,或者扩散模型。扩散模型效果好但推理慢,先在服务器上验证,再决定要不要上。

从工程落地看,大多数下载即用的 python 源码包中,模型入口函数一般长这样:

def forward(self, image, mask): x = torch.cat([image, mask], dim=1) # 通道维拼接 x = self.encoder(x) x = self.decoder(x) return x

其中 image 的 shape 是(B, 3, H, W),mask 是(B, 1, H, W)。这个接口决定了输入是归一化后的 float 还是 0~255 的 int。如果直接把 uint8 的 mask 送进去,模型第一层卷积会把它当成大数值特征,训练初期梯度就崩了。正确做法是mask = mask.float() / 255.0,让值域落在 [0, 1]。

3.2 分割 + 条件修复的两段式

如果手写笔迹的形态比较复杂(细、浅、与背景重叠),一步修复的方式往往会在「定位」上浪费大量能力。两段式的思路是:

第一段:训练一个分割网络(轻量 U-Net 或 HRNet)预测手写区域的概率图;第二段:把概率图当作 mask 输入修复网络。这样做的优点是把「找手写」和「擦手写」解耦,可以单独调试;缺点是错误传播,第一段漏检的地方第二段不会主动补。

一个工程折中是用迭代细化的方式:把分割输出的概率图做 softmax,不硬阈值,作为修复网络的输入。修复网络可以自己学会「概率低的地方少动、概率高的地方多做填充」。在扫描件质量差的数据集上,soft 概率图比 hard mask 平均高 1.5~2dB PSNR。

3.3 用 PyTorch 搭建最小可跑的门控卷积骨架

下面给一个门控卷积的最小实现,保证在单卡或 CPU 上能跑通,用来验证数据和 loss 流程。

import torch import torch.nn as nn class GatedConvBlock(nn.Module): def __init__(self, in_c, out_c, k=3, s=1, p=1): super().__init__() self.conv = nn.Conv2d(in_c, out_c, k, s, p) self.gate = nn.Conv2d(in_c, out_c, k, s, p) self.relu = nn.ReLU(inplace=True) self.sigmoid = nn.Sigmoid() def forward(self, x): feature = self.relu(self.conv(x)) gate = self.sigmoid(self.gate(x)) return feature * gate

这个 block 的核心是两条并行卷积支路:feature 支路提取语义特征,gate 支路输出一个 0 到 1 之间的门控系数,逐元素相乘后决定特征保留多少。门控让网络在 mask 边缘自动学习平滑过渡,比普通卷积对边界更友好。把它拼成编码器解码器时,下采样用 stride=2 卷积,上采样用双线性插值或转置卷积,中间加一个 bottleneck。最小模型:encoder 三层、decoder 三层,约 3MB 参数,CPU 上也能推理一张 512x512 图。

跑通这个 baseline 后再考虑换 backone。不建议一上来就上扩散模型,数据量不大时,带门控卷积的 U-Net 在稳定性和收敛速度上都更占优。

4. 训练配置:损失函数组合与超参数

4.1 四个常用损失及其权重范围

手写擦除不是简单的像素回归任务,单独用 L1 会产生模糊结果。我的默认组合是:

损失计算方式作用权重建议
mask 内 L1只在 mask 覆盖区域算强制擦除区域重建10
mask 外 L1在全图除 mask 外区域算保护背景不被改动1
感知损失VGG16 conv1_2..conv4_2 特征 L1保留结构与语义0.1
对抗损失PatchGAN 判别器让纹理更真实0.05

mask 内 L1 权重最大,因为被遮住的地方完全没有像素监督,给大权重才能逼着生成器学会填内容;mask 外 L1 权重小但非常重要,它约束模型不要越界修改。感知损失帮助保留字形和扫描纹理的高频,对抗损失只给一个很小的系数,防止早期噪声主导训练。

用 PyTorch 组合多 loss 的常见写法如下:

import torch import torch.nn as nn import torch.nn.functional as F # vgg_layers 是预提取的 VGG16 中间层特征,这里省略实现 def vgg_loss(pred, target): feat_p = vgg_layers(pred) feat_t = vgg_layers(target) return sum(F.l1_loss(a, b) for a, b in zip(feat_p, feat_t)) def train_step(model, disc, opt_g, opt_d, image, mask, target): # image/target 转成 [-1,1] 或 [0,1],mask 转成 float [0,1] x = torch.cat([image, mask], dim=1) pred = model(x) # 1) L1 loss,mask 内外分权重 diff = torch.abs(pred - target) loss_l1 = (diff * mask).mean() * 10 + (diff * (1 - mask)).mean() * 1 # 2) 感知损失 loss_p = vgg_loss(pred, target) * 0.1 # 3) generator 的对抗损失,用 relu 形式近似 hinge fake_logit = disc(pred) loss_g = -fake_logit.mean() * 0.05 loss = loss_l1 + loss_p + loss_g opt_g.zero_grad() loss.backward() opt_g.step() # 4) 判别器更新 real_logit = disc(target) fake_logit = disc(pred.detach()) loss_d = (F.relu(1 - real_logit) + F.relu(1 + fake_logit)).mean() opt_d.zero_grad() loss_d.backward() opt_d.step()

mask 在这里有两个用途:输入里作为额外通道,loss 里作为权重图。两者共用同一个 mask 张量,但建议在 loss 计算时对 mask 做一次轻微膨胀或高斯平滑,避免在 mask 边界处出现剧变的像素级损失跳跃。判别器更新用的是 hinge loss 而不是 BCE,这是 GAN 训练里稳定性更高的小 trick。

4.2 训练参数与 schedule

从实测中沉淀下来的一套比较稳的配置:

参数推荐值说明
分辨率256x256 训练,512x512 微调先低后高,收敛快
batch size8~16看显存,小卡用 8
优化器Adam,lr=0.0002,beta=(0.5, 0.999)GAN 常用
训练比例G 两步,D 一步防止判别器过强
迭代步数初始 100k,微调 50k观察验证 loss

这套配置适合大多数基于 GAN 的手写擦除数据集。如果用的是纯 L1 + 感知损失、不加 GAN,学习率可以提到 0.0004,训练更稳。

4.3 三个常见训练失败模式的排查

失败模式一:loss 在降但输出模糊。原因是感知损失权重太小,L1 主导,生成器学会了「平均化」内容。调法是把感知损失权重升到 0.2~0.3,同时把对抗损失降到 0.02。

失败模式二:输出边缘发黑或出现蓝绿色伪影。大概率是目标张量没归一化对:如果输出层用 Tanh,值域是 [-1,1],target 也要同步到 [-1,1];用 Sigmoid 则 target 要在 [0,1]。混用后模型再怎么学习都输出受限。

失败模式三:整张图被重绘。通常是训练时 crop 的 mask 占比过高(超过 40%),模型损失被 mask 内部主导,外部约束不起作用。解决方法是采样时限制 mask 占比,按 2.4 里说的 10%~30% 执行。

提示:排查这些问题时,先看训练日志里第一个 batch 的 image、mask、target、pred 四张图的可视化对比,比看 loss 曲线更直接。

5. 推理与效果评估:PSNR 不是唯一指标

5.1 模型导出与批量推理脚本

训练好的权重,在 python 包里的标准做法是 torch.load 后 model.eval()。实际工程里,我更建议先转成 ONNX,后续部署和验证都方便。

import torch from model import build_model model = build_model(weights='best.pth') model.eval() dummy_img = torch.randn(1, 3, 512, 512) dummy_mask = torch.zeros(1, 1, 512, 512) torch.onnx.export( model, (dummy_img, dummy_mask), "eraser.onnx", opset_version=17, input_names=['image', 'mask'], output_names=['output'], dynamic_axes={'image': {0: 'batch'}, 'mask': {0: 'batch'}} )

这里 dynamic_axes 只把 batch 设为动态,宽高保持固定 512x512。这样导出的 ONNX 在转 TensorRT 时不会有动态 shape 带来的性能损失。如果有多个分辨率需求,按每个分辨率单独导出,而不是用一个动态 H/W。

ONNX 导出后用 onnxruntime 验证输出和 PyTorch 的差异:

import onnxruntime as ort import numpy as np sess = ort.InferenceSession("eraser.onnx", providers=['CPUExecutionProvider']) out = sess.run(None, { 'image': img_np.astype(np.float32), 'mask': mask_np.astype(np.float32), })[0]

这个步骤主要查两件事:输入张量的通道顺序(NCHW 还是 NHWC)是否符合 runtime 期望;归一化是否已被包含进模型。如果模型输入是[1, 3, 512, 512]的 NCHW,前处理要先做 HWC 转 NCHW,再做归一化并转 float32。很多推理错误最后都归结到输入格式上——模型本身没错,是预处理没做全。

5.2 评估指标怎么选

学术指标上,手写擦除报告通常会同时给 PSNR 和 SSIM,但我更看重 LPIPS 和业务端的漏擦率。下面是我常用的一套评估口径:

指标关注点使用说明
PSNR像素级差异易受背景平滑影响
SSIM结构相似度对边缘敏感
LPIPS感知语义更接近人眼
漏擦率手写残留业务判定

漏擦率的计算方式:把模型预测结果和真实 clean 图做差,差值超过阈值的像素数除以 mask 内像素总数。threshold 通常取绝对差值 > 50(针对 0~255 灰度图)。

def calc_removal_metrics(pred, target, mask): # 先归一化到 0~1,再乘 255 统一量纲 pred = (pred - pred.min()) / (pred.max() - pred.min() + 1e-8) diff = (pred - target).abs().mean(dim=1, keepdim=True) * 255 residual = (diff > 50).float() * mask missed = residual.sum() / (mask.sum() + 1e-8) return missed.item()

这里把 pred 和 target 归一化到 0~1 再乘 255,是为了统一不同预处理差异。missed 接近 0 说明擦得干净,但要注意 missed 低不等于效果好:如果模型把背景也重绘了,目标差值不满足阈值,missed 仍会低。所以最终还要配合 LPIPS 或目检。

6. 部署环节的 3 个提速技巧

6.1 把预处理和后处理合并进 ONNX

常见做法是在前处理阶段把归一化、ToTensor、HWC 转 NCHW 写成 numpy 矩阵操作,但部署时这些操作会占用大量 CPU 时间——尤其是一张 512x512 的图要经历 uint8 转 float、permute、除法三个步骤。一种更省事的做法是把归一化直接写进 ONNX 图里:给模型包一层 wrapper,输入原图 uint8,输出模型所需的 float。

class WrappedModel(nn.Module): def __init__(self, inner): super().__init__() self.inner = inner def forward(self, img_uint8, mask_uint8): img = img_uint8.float() / 255.0 mask = mask_uint8.float() / 255.0 img = img.permute(0, 3, 1, 2) mask = mask.permute(0, 3, 1, 2) return self.inner(img, mask)

保存后再转 ONNX,输出端直接拿干净图。注意:用 onnxruntime 验证时,输入 dtype 已变为 uint8,不要再重复做归一化。

6.2 把 mask 膨胀放进批量推理流程

训练时你使用膨胀后的 mask 做 loss,但推理时从分割网络出来的 mask 往往是未膨胀的,会导致擦除范围偏小、残留毛边。我一般会在推理链路最后加一次 cv2.dilate,用 5x5 椭圆核迭代 1 次,把 mask 边缘往外扩 2 像素。这个操作如果写在 Python 循环里逐个处理,会很慢;对整批 mask 用 numpy 一次做完会快很多:

import cv2 import numpy as np # m: shape (N, 1, H, W) 的 uint8 mask mask_np = m.cpu().numpy().transpose(0, 2, 3, 1) # NCHW -> NHWC kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (5, 5)) dilated = np.stack([ cv2.dilate(mask_np[i], kernel, iterations=1) for i in range(mask_np.shape[0]) ]) mask = torch.from_numpy(dilated).permute(0, 3, 1, 2).to(m.device)

这样单 batch 开销可以忽略,也避免了在 GPU 数据流里做 OpenCV 操作带来的同步开销。

6.3 半精度推理与 CPU 预处理解耦

如果模型保持 PyTorch 推理(不转 ONNX),可以把模型和数据都切到半精度。手写擦除这类像素级输出对精度不敏感,FP16 输出通常肉眼不可查,做法是model.half()data.half()。注意:如果后续还要用 torch.compile,先切 half 再 compile,避免编译阶段锁定 fp32 卷积参数。

至于并发,最容易见效的是把 CPU 预处理(读图、缩放、mask 膨胀)放到多进程队列里,GPU 侧只做模型前向。Python 的多线程无法真正并行 CPU 计算,用 multiprocessing pool 做数据加载和 mask 膨胀,配合 queue 喂给推理进程,吞吐量能提升 40% 左右。手写擦除这类任务对延迟敏感度不高,通常跑在批量队列里,瓶颈多半在 CPU 预处理而不是 GPU 算力——先排查这个方向,而不是一味换大模型。

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

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

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

立即咨询