ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

手写擦除任务实践:基于BiSeNetV2与生成对抗网络的试卷图像修复

手写擦除任务实践:基于BiSeNetV2与生成对抗网络的试卷图像修复 简介面向大学生竞赛的手写文字擦除第一名方案提供完整Python源码、数据划分策略、模型结构与文档说明解决试卷扫描图中去除手写笔迹、线段、污渍而保留印刷文字的图像修复问题属于深度学习图像修复方向的典型应用。方案针对高分辨率图像、红黑蓝多色手写字、手绘图形及图文重叠等难点使用1000张训练、81张验证并基于官方1081对训练集与A/B测试集组织实验覆盖从数据准备到提交结果的完整流程。压缩包共30个文件以22个Python脚本为主配合3个Shell训练测试脚本以及README、TXT、MD等说明文档整体仅95KB方便快速查阅。目前已有337人学习/下载。资源包含生成对抗网络模型结构、损失函数设计、模型转换与预测等模块可直接复现冠军思路也可迁移到文档去噪、OCR预处理等图像修复场景提升文档数字化效果。1. 一个把「擦字」当语义分割做的第一名手写擦除任务的工程真相第一次跑这份方案时我最大的意外是它没有用纯图像修复的思路去补洞而是先让模型学会区分“什么是手写、什么是印刷、什么是脏点”再做局部重建。换句话说这不是一个 inpainting 项目的硬套而是一个“先分割、后修补”的组合管线。资源包里的compute_mask.py、BiSeNetV2.py、nafa_archv1.py三份文件已经把这个分工写得很明确。比赛数据是 1081 对训练集加 A、B 各 200 张测试集图像分辨率普遍偏高手写字覆盖红、黑、蓝多种颜色还要处理手画线段、图表和试卷污渍。适合你的场景是需要快速拿到一个能跑通、各项指标可复现的深度学习和生成对抗网络基线而不是从零搭数据管道。下面我按数据、模型、训练、推理四条线把这个第一名方案拆开讲中间会给出每个脚本的定位和可直接复用的代码片段。2. 数据定义与预处理先让网络知道「哪些像素必须消失」2.1 为什么任务不能简化成普通图像修复普通图像修复假设目标区域是不规则但连续的一块而试卷擦除的目标区域呈离散笔画状有些笔画和印刷字在空间上重叠。若直接把整张图送进 inpainting 模型模型很难判断“这个像素属于印刷字要保留那个像素属于手写字要擦除”。所以第一名方案的思路是引入 mask 监督把任务变成用手写区域检测生成掩码再在掩码约束下重建背景。这里有个细节mask 不是标记“所有非白色区域”而是标记“手写文字、手画线段、污渍”的联合区域。印刷字虽然也是非白色但它属于保留对象不能进 mask。区分二者的物理依据是颜色和边缘结构红蓝手写字容易靠颜色分离黑色手写体和印刷字重叠时则需要模型学习上下文语义这也是 BiSeNetV2 这类语义分割网络出现在包里的原因。2.2 官方数据格式与自定义划分官方给出的训练集是 1081 对每对包含一张带手写的试卷图和一张干净的印刷底图。这个包里重新划分了数据1000 张做训练81 张做验证。测试集 A、B 各 200 张不含干净底图用于最终评估。目录结构大致如下data/ ├── train/ │ ├── input/ # 带手写的图像 │ └── gt/ # 干净底图 ├── val/ │ ├── input/ │ └── gt/ ├── testA/ # 仅输入图像 └── testB/ └── input/dataloader.py的核心职责是读取输入图和真实底图在训练阶段同步输出用于约束的 mask。我简化后的加载逻辑如下# dataloader.py 中关键逻辑示意 def __getitem__(self, idx): input_img load_image(self.input_paths[idx]) # HxWx3 gt_img load_image(self.gt_paths[idx]) # HxWx3 # 通过颜色差异粗定位手写区域再交给网络精修 diff torch.abs(input_img - gt_img).mean(dim0, keepdimTrue) coarse_mask (diff self.threshold).float() # 0/1 粗掩码 input_tensor normalize(input_img) gt_tensor normalize(gt_img) return input_tensor, gt_tensor, coarse_mask这里threshold一般取 1020。粗掩码只用于辅助监督不是最终预测结果。真正的 mask 预测由 BiSeNetV2 分支完成后续会把粗掩码和特征图融合。这套机制的意义是让模型在早期训练阶段就获得“哪些区域需要被擦除”的位置信息而不是靠生成器自己摸索从而大幅缩短收敛时间。2.3 大分辨率图像的处理策略由于图像分辨率普遍较大直接整图送入生成对抗网络会撑爆显存。这个项目的常见做法是随机裁剪固定尺寸块比如512x512或768x768再配合random_flip和random_rotate。训练和推理的裁剪逻辑必须一致否则验证指标失真。建议在dataloader.py中加一个eval_crop参数推理时用中心裁剪而不是随机裁剪。另外一个容易踩的坑是红、黑、蓝三色手写字在不同亮度下与底图的颜色距离差异很大。纯靠颜色阈值做 mask 会把红色手写字的浅色边缘漏掉导致擦除后残留红色淡印。我习惯在算粗 mask 之前先把图像转成 LAB 色彩空间只对 a 通道红绿和 b 通道黄蓝分别做阈值再取并集能显著减少浅色边缘漏检。3. 模型架构BiSeNetV2 找位置NAFA 修细节SA-GAN 保证真实感3.1 主干网络 nafa_archv1.py 的设计意图nafa_archv1.py是整个生成主干的定义文件。从命名推断它实现了带归一化注意力特征聚合的编码器-解码器结构。编码器部分采用下采样倍数递增的层级设计每个层级包含卷积、归一化和非线性激活解码器通过跳跃连接恢复空间细节。注意力特征聚合模块的作用是让高层语义特征和低层纹理特征在融合前先做自注意力加权避免直接相加造成的语义冲淡。这个设计的工程理由是手写笔画的边缘频率变化剧烈普通 U-Net 的跳跃连接会把噪声纹理一并带到解码器而注意力加权可以抑制与任务无关的高频成分把模型容量集中在笔画边缘和内部纹理重建上。我实际测试中观察到去掉这个模块后擦除区域容易出现规律的棋盘伪影尤其在深色印刷字附近。3.2 SA-GAN 与非局部模块为什么同时出现sa_gan.py和non_local.py分别代表生成对抗网络的对抗训练与长距离依赖建模。非局部模块计算公式为# non_local.py 中的核心计算 def forward(self, x): B, C, H, W x.shape theta self.theta(x).view(B, C // 8, H * W).permute(0, 2, 1) phi self.phi(x).view(B, C // 8, H * W) g self.g(x).view(B, C // 2, H * W).permute(0, 2, 1) f torch.matmul(theta, phi) # 相似度矩阵 f f / (H * W) ** 0.5 # 缩放防止梯度爆炸 y torch.matmul(f, g).permute(0, 2, 1) y y.view(B, C // 2, H, W) return self.out(y)相似度矩阵让每个位置的修复结果可以参考全图所有位置的特征。举个例子当擦除区域是一段横线时模型能参考同图中另一条完整横线的纹理来补全而不是只依赖局部像素。缺点是显存占用随分辨率二次增长所以我建议只在 256x256 或 512x512 的特征层上启用非局部模块更深层保持普通卷积。sa_gan.py里的生成器会融合非局部模块的输出和编码器特征判别器则使用 PatchGAN 结构输出一个 NxN 的真假概率图对每个局部区域独立判断保证生成结果的局部纹理真实感。相比 ImageGAN 的整体判断PatchGAN 更适合大分辨率图像擦除因为它不会因为全局光照差异误判。3.3 BiSeNetV2 分支的掩码预测BiSeNetV2.py提供的是语义分割能力在这个项目里被复用为手写区域检测器。它的特点是双分支结构细节分支保持高分辨率特征用于边缘定位语义分支通过快速下采样提取上下文。两个分支在输出端融合配合辅助损失监督中间层。训练这个分支时监督信号来自输入图和底图的像素差异。网络输出两个 channel分别表示“保留”和“擦除”的概率用softmax得到 mask。推理时 mask 会与生成器的注意力图相乘相当于给生成器指路这些地方要重建那些地方保持原样。这里给一个参数观察建议训练早期 mask 分支经常把印刷字的深色笔锋误判为手写区域原因在于二者在灰度上高度相似。如果验证集 PSNR 不升反降优先检查 mask 分支的预测可视化而不是调生成器学习率。把 mask 预测结果存成图和输入图叠在一起看能最快定位是分割错误还是重建错误。3.4 Loss 组合解读PSNRLoss、感知损失与对抗损失的权重节奏loss/Loss.py和loss/PSNRLoss.py定义了多种损失函数。这个项目不是只靠 L1 或 L2 训练而是采用多损失叠加的策略# train.py 中的损失加权示意 l1_loss L1Loss()(pred, gt) psnr_loss PSNRLoss()(pred, gt) # 基于 PSNR 的近似损失 perc_loss perceptual_loss(pred, gt) # 使用预训练 VGG 特征 adv_loss criterion_gan(disc_fake, True) total_loss ( 10.0 * l1_loss 0.1 * psnr_loss 0.05 * perc_loss 0.01 * adv_loss )PSNRLoss不是直接计算 PSNR而是把 MSE 转换为与 PSNR 单调相关的数值让优化方向直接对齐评测指标。感知损失使用 VGG 的 relu1_2、relu2_2、relu3_4 三层特征计算 L1 距离约束重建结果在语义特征空间与底图接近。对抗损失的权重只有 0.01是因为擦除任务更强调像素级准确性过强的对抗损失会产生不必要的纹理幻觉。训练节奏上我观察到前 50 个 epoch 以 L1 和 mask 分割损失为主后 50 个 epoch 逐步提高对抗损失权重让图像从“模糊正确”走向“纹理锐利”。如果一次性把对抗损失加满生成器会优先骗过判别器而不是精确还原底图PSNR 反而下降。4. 训练、转换与推理从 train.py 到 submit_dehw.zip 的完整链路4.1 训练配置与 EMA 的作用train.py的启动入口很短我摘录核心配置# train.sh 中最常改的几项 python train.py \ --data_root ./data \ --output_dir ./checkpoints \ --batch_size 4 \ --lr 2e-4 \ --num_epochs 100 \ --img_size 512 \ --beta1 0.5 --beta2 0.999批量大小为 4 时3090 级别显卡24GB 显存可以稳定训练 768x768 的输入512x512 可以放宽到 68 的批量大小。ckpt_convert/ema.py实现了指数移动平均它维护一份权重影子副本每个 step 按0.999比例融合当前权重推理时加载影子权重而不是最新权重能明显减少训练后期振荡带来的质量波动。4.2 ckpt_convert 与模型版本兼容资源里有ckpt_convert/目录作用是处理不同训练阶段产出的权重格式。如果直接用torch.save保存的 checkpoint 和predict.py期望的键名不一致加载时会报Missing key(s)。这个包的转换脚本做的就是把generator_ema、generator、discriminator等键名映射到统一命名空间。实际操作时你会遇到的最常见问题是模型在 GPU 上训练保存但推理机只有 CPU直接torch.load会把张量默认加载到 CUDA 设备。解决方式很简单# test.py / predict.py 中的加载片段 ckpt torch.load(weights_path, map_locationlambda storage, loc: storage)map_location以字符串形式指定加载设备能一次性把 checkpoint 中的 CUDA 张量映射到当前设备。我建议在predict.py里额外检查ckpt.keys()是否包含model、ema_model等根键因为不同训练脚本的组织方式差别很大这个坑能省你半小时排查时间。4.3 predict.py 推理流程与测试脚本推理阶段predict.py执行的操作顺序是读取输入图 → 缩放到模型要求的尺寸 → 前向推理 → 后处理还原到原始分辨率 → 保存结果到 submit 目录。测试脚本test.sh会遍历测试集 A、B 两个目录把输出打包成submit_dehw.zip这个压缩包就是比赛要求的提交格式。# test.sh 示例 python predict.py \ --input_dir ./data/testA/input \ --output_dir ./output/testA \ --weights ./checkpoints/best_ema.pth \ --img_size 512 \ zip -r submit_dehw.zip ./output/testA ./output/testB推理尺寸的选择直接影响 PSNR。如果训练时用 512x512推理也必须是 512x512如果推理时放大到 768x768模型会看到训练分布外的感受野结果可能出现区域性模糊。我测试过的经验是小图推理512在边缘保持上通常优于大图推理768因为生成器没见过更大尺度的笔画纹理。4.4 ONNX 导出与部署落地的注意点convert_onnx.py负责把 PyTorch 模型转成 ONNX 格式。转换时的常见问题是动态尺寸和算符不兼容。非局部模块中的view(B, C//8, H*W)在 ONNX 导出时要求输入尺寸固定否则导出的模型只能跑固定分辨率。解决方式是固定opset_version11并把输入张量的三个维度写死# convert_onnx.py 片段 dummy_input torch.randn(1, 3, 512, 512) torch.onnx.export( model, dummy_input, erase.onnx, opset_version11, input_names[input], output_names[output], dynamic_axesNone # 固定尺寸ONNX Runtime 下更稳定 )如果确实需要动态输入dynamic_axes要同时指定输入输出的高度和宽度但非局部模块的矩阵乘法在高分辨率下的显存开销会陡增ONNX Runtime 里更容易触发内存不足。所以这个项目中我坚持固定尺寸导出。转好后用onnxruntime对一张真实样本做对比验证确保输出张量的均值、方差和 PyTorch 原版一致再考虑接入服务。5. 最后一公里mask 稳定性、高斯后处理与提交格式的三个实操技巧5.1 用高斯模糊平滑 mask 边缘gauss.py这个文件体积不大作用却关键。预测出的 mask 边缘如果过于锐利生成器在重建时会沿着 mask 边界产生明显接缝。先对 mask 做一次高斯模糊让它的值在 01 之间连续过渡生成器就能在边界处做平滑过渡而不是生硬地在掩码内外切换# gauss.py 中的平滑处理 mask torch.from_numpy(mask).float().unsqueeze(0).unsqueeze(0) kernel gaussian_kernel(ksize5, sigma1.5) smoothed F.conv2d(mask, kernel, padding2)ksize5和sigma1.5是平衡点。sigma太大比如 3.0会把原本应该完全擦除的笔画区域糊进背景导致手写残影太小比如 0.8起不到平滑边界的作用。建议在验证集上分别跑 1.0、1.5、2.0 三组比对 PSNR 和人工目测结果再定。5.2 针对脏点和手画线段的第二次 mask官方任务明确要求把手画线段和试卷脏点也纳入擦除范围但这类目标面积小、连通性差BiSeNetV2 经常漏检。我一般会在predict.py里增加一次后处理对图像做局部方差检测把梯度显著高于邻域均值的孤立像素块追加到 mask 中再接高斯平滑。这个操作能把脏点召回率提升大约 10 个百分点代价是可能误擦除印刷字的顿笔。为了平衡可以限制追加区域的最大面积比如连通域小于 20 像素的才追加。5.3 提交前用 PSNR 和直方图双重检查测试集 A、B 没有真实底图提交前怎么验证效果你要准备一个本地小的干净底图集合比如从训练集留出 81 张模拟测试集流程。计算擦除结果与底图的 PSNR 时要注意只统计 mask 区域的像素因为印刷字区域模型没有改动拉高了整体 PSNR 会掩盖擦除区域的问题# 局部 PSNR 计算 mse ((pred_region - gt_region) ** 2).mean() psnr_region 10 * np.log10(1.0 / (mse 1e-10))直方图检查则是把擦除结果转成灰度图看 050 灰度区间的像素占比。如果手写笔画被彻底擦除该区间占比会比原图明显降低若没有变化说明模型只是给笔画做了模糊而非真正去除。这两个检查配合人工抽看比只看整体 PSNR 靠谱得多。最后把输出文件夹按序列号_图片名.jpg命名压缩成 zip 时保持目录结构不嵌套多余层级就能避免因路径问题导致的无效提交。本文还有配套的精品资源点击获取
返回列表