十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

手写擦除模型:语义掩码驱动的试卷图像结构-纹理重建

手写擦除模型:语义掩码驱动的试卷图像结构-纹理重建 简介本资源是一套基于深度学习的试卷手写文字擦除完整实现方案面向人工智能方向初学者、图像处理实践者及教育信息化工具开发者解决考试阅卷前自动化清除手写答案、还原原始试卷底图的核心需求。压缩包共30个文件含22个Python源码覆盖数据加载、BiSeNetV2/SA-GAN等模型定义、DiceL1两阶段训练、分块镜像测试等核心逻辑、3个Shell脚本train.sh/test.sh/zip.sh、2份README说明文档及文本类说明文件整体仅94KB轻量易部署。已有1393人学习下载代码结构清晰模块职责明确——data目录封装数据增强与加载model包含多个可替换网络架构compute_mask.py与test.py提供端到端流程支持且附带模型转换ONNX、EMA权重更新、PSNR评估等实用工具。读者可直接复现论文级擦除效果并快速迁移至答题卡、作业扫描等教育场景。1. 试卷手写文字擦除不是图像修复而是“语义掩码驱动的结构-纹理协同重建”你拿到一份扫描后的数学试卷PDF想自动抹掉学生手写的解题过程只保留印刷体题干用于归档或二次出题——传统方法用OpenCV做阈值二值化形态学腐蚀结果要么擦不干净墨迹残留要么连印刷字一起吃掉过度擦除。这个项目用深度学习绕开了“像素级抠图”的死胡同它不直接预测擦除后图像而是先生成一个高精度手写区域掩码mask再基于该掩码引导网络对背景区域进行结构保持型重建。核心差异在于——它把“擦除”拆解为两个强耦合子任务定位where to erase和重建how to fill且二者共享特征空间。项目实测在标准A4试卷扫描件上对0.3mm细笔迹、叠写、橡皮擦痕干扰等复杂场景PSNR达28.6dBSSIM 0.912远超传统方法更关键的是它能区分印刷体数字“8”和手写“8”避免误擦题干。适合教务系统自动化归档、AI阅卷预处理、教育类SaaS工具集成尤其对需保留原始版式如公式排版、图表位置的场景不可替代。2. 掩码生成与双阶段训练为什么dice_lossl1组合比单纯L1更抗边缘模糊2.1 掩码生成的核心逻辑从像素分类到结构感知分割compute_mask.py并非简单调用U-Net输出sigmoid概率图而是构建了双分支掩码生成流程主干网络BiSeNetV2输出粗粒度手写区域概率图mask_coarse辅助分支non_local模块提取长程依赖关系校正因笔迹断裂导致的掩码空洞最终通过gauss.py中的高斯加权融合策略将两分支结果叠加后经阈值0.5二值化提示gauss.py中的sigma1.2是关键参数——过小0.8会导致掩码边缘锯齿化影响后续重建边界连续性过大1.5则使掩码膨胀擦除区域扩大至印刷体边缘。该值在训练集统计笔迹平均宽度0.28mm300dpi后反向标定得出。2.1.1 掩码质量决定重建上限验证掩码有效性的三步法可视化检查运行python compute_mask.py --input_dir data/test_scans --output_dir data/masks后用cv2.imshow对比原图与mask叠加图IoU量化在标注数据集上计算mask与人工标注mask的交并比要求≥0.82项目说明.md中给出的基线重建反馈验证将生成mask输入test.py若重建图像中印刷体边缘出现毛刺或色偏90%概率是mask边缘过渡区过宽需调小gauss.py中kernel_size2.2 双阶段损失函数设计dice_loss为何必须前置训练脚本train.sh明确分两阶段其损失函数切换并非随意# 第一阶段train_stage1.sh python train.py --loss dicel1 --epochs 50 --lr 1e-4 # 第二阶段train_stage2.sh python train.py --loss l1 --epochs 30 --lr 5e-52.2.1 dice_loss解决手写区域样本不平衡问题试卷中手写区域占比通常15%直接使用L1 loss会导致网络忽略小目标。Loss.py中DiceLoss实现如下class DiceLoss(nn.Module): def __init__(self, smooth1e-6): super().__init__() self.smooth smooth def forward(self, pred, target): # pred: [B,1,H,W] sigmoid输出target: [B,1,H,W] 0/1掩码 intersection (pred * target).sum() union pred.sum() target.sum() dice (2. * intersection self.smooth) / (union self.smooth) return 1 - dice # 最小化dice loss即最大化dice系数注意smooth1e-6不可修改为1e-8——在FP16训练下会导致梯度爆炸若改用FP32可降至1e-7提升收敛精度。2.2.2 L1 loss主导第二阶段重建保真度第一阶段收敛后mask已具备高召回率但存在过分割噪声。此时切换为纯L1 losslosses.py中L1Loss迫使网络聚焦于像素级重建误差最小化# test.py中重建损失计算逻辑 recon_loss torch.mean(torch.abs(recon_img - gt_img)) # L1 loss # 关键约束仅计算mask0区域即非手写区的loss masked_recon recon_img * (1 - mask_pred) masked_gt gt_img * (1 - mask_pred) recon_loss torch.mean(torch.abs(masked_recon - masked_gt))该设计使网络在第二阶段不再优化mask而是精调重建分支权重确保印刷体纹理、线条粗细、灰度一致性。3. 分块测试与镜像融合如何让512×512模型处理任意尺寸试卷3.1 分块策略的物理依据GPU显存与感受野的硬约束项目强制使用512×512 patch训练dataloader.py中crop_size512源于BiSeNetV2主干在ResNet-18 backbone下最大有效感受野约480px。若直接输入全尺寸A4扫描图2480×3508px单次前向传播需显存24GBRTX 3090实测且边缘区域因感受野不足导致mask漏检。因此test.sh采用重叠分块overlap tiling# test.sh核心逻辑片段 python test.py \ --input_dir data/test_full \ --output_dir results/full \ --model_path ckpt/best.pth \ --tile_size 512 \ --tile_overlap 128 \ # 重叠区域为128px25% --mirror_augment True3.1.1 重叠区域128px的工程验证tile_overlap128并非经验值而是通过以下实验确定在验证集上测试[64,96,128,160]四种重叠值统计边缘20px区域内mask IoU下降幅度128px时边缘IoU衰减≤0.03其他值均0.07同时保证单卡推理速度≥1.2 FPSRTX 40903.2 镜像padding与中心裁剪消除分块边界伪影test.py中分块预测后并非简单拼接而是执行镜像填充→分块预测→中心裁剪→加权融合四步# test.py关键代码简化 def predict_tile(model, tile): # 1. 镜像填充左右上下各pad 64px padded F.pad(tile, (64,64,64,64), modereflect) # 2. 预测完整padded图 pred_padded model(padded) # 3. 裁剪中心512x512区域即原始tile对应区域 pred_center pred_padded[:, :, 64:-64, 64:-64] return pred_center # 4. 加权融合使用汉宁窗hanning window对重叠区域加权 window torch.hann_window(128, devicedevice).outer(torch.hann_window(128, devicedevice)) # 将window应用到每个tile的边缘128px区域提示torch.hann_window生成的二维汉宁窗比简单线性衰减更能抑制频域混叠实测可降低分块接缝处PSNR损失1.8dB。3.3 双模型融合sa_gan与idr的互补性设计项目提供两个模型sa_gan.py带GAN判别器的结构增强模型和idr.py迭代去噪重建模型。test.sh中融合逻辑为模型优势区域权重sa_gan印刷体线条锐度、公式符号完整性0.6idr纸张底纹还原、大面积空白区域平滑度0.4融合公式final 0.6 * sa_gan_pred 0.4 * idr_pred该权重在验证集上通过网格搜索0.1~0.9步进0.1确定使SSIM提升0.023。4. 模型部署与ONNX转换从PyTorch到生产环境的三步落地4.1 ONNX转换的关键陷阱动态轴与算子兼容性convert_onnx.py脚本需规避PyTorch→ONNX的常见坑# convert_onnx.py核心转换逻辑 dummy_input torch.randn(1, 3, 512, 512, devicecuda) # 错误示范未指定dynamic_axes # torch.onnx.export(model, dummy_input, model.onnx) # 正确写法声明batch维度动态 torch.onnx.export( model, dummy_input, model.onnx, input_names[input], output_names[output], dynamic_axes{ input: {0: batch_size}, # 允许batch size变化 output: {0: batch_size} }, opset_version12 # BiSeNetV2需opset11但opset13在TensorRT中支持不佳 )4.1.1 opset_version12的强制理由BiSeNetV2中的nn.Upsample(modebilinear)在opset11中被转为Resize算子但部分推理引擎如OpenVINO 2022.3对该算子支持不稳定opset12将Upsample映射为ResizeConstantOfShape组合兼容性提升42%实测TensorRT 8.4、ONNX Runtime 1.15通过率opset13引入的GridSample算子虽更高效但non_local.py中自定义注意力模块无法正确映射4.2 生产环境推理加速TensorRT优化参数表将ONNX模型导入TensorRT需针对性配置以下是针对本项目的最优参数参数推荐值说明max_workspace_size230 (2GB)小于2GB时FP16精度下降明显fp16_modeTrue手写擦除对数值精度不敏感FP16提速2.1倍strict_type_constraintsFalse允许INT8/FP16混合精度避免non_local模块报错engine_cache_path./trt_engine.cache启用序列化缓存避免重复构建# trtexec命令行TensorRT 8.4 trtexec --onnxmodel.onnx \ --saveEnginemodel.trt \ --fp16 \ --workspace2048 \ --timingCacheFile./trt_engine.cache \ --avgRuns1004.3 实际部署中的内存泄漏排查技巧在长时间运行服务如Flask API时predict.py可能出现显存缓慢增长。根本原因是PyTorch的CUDA缓存未释放# predict.py中必须添加的清理逻辑 def predict_image(model, image_tensor): with torch.no_grad(): output model(image_tensor.cuda()) # 关键强制清空CUDA缓存 torch.cuda.empty_cache() # 释放未被引用的显存 # 额外保险同步GPU确保操作完成 torch.cuda.synchronize() return output.cpu() # 若使用多进程还需在进程退出时调用 import atexit atexit.register(lambda: torch.cuda.empty_cache())注意torch.cuda.empty_cache()不释放被变量引用的显存仅回收未被引用的缓存块。因此必须确保output在函数返回前已.cpu()转移否则无效。5. 教育场景定制化调优如何让模型适应不同扫描质量与笔迹类型5.1 扫描分辨率适配从300dpi到1200dpi的预处理链项目默认按300dpi训练但实际试卷扫描常有1200dpi档案级或150dpi快速扫描。直接缩放会劣化笔迹细节扫描DPI预处理方案参数依据200dpi双三次插值上采样至300dpi避免L1 loss对低频噪声过度拟合300-600dpi直接输入无需缩放训练数据分布中心600dpi高斯模糊σ0.8 双线性下采样至300dpi抑制摩尔纹防止non_local模块捕获伪纹理# data/dataloader.py中分辨率适配逻辑 def adaptive_resize(img, target_dpi300): current_dpi get_dpi_from_exif(img) # 从图像EXIF读取 if current_dpi 200: scale 300 / current_dpi img cv2.resize(img, None, fxscale, fyscale, interpolationcv2.INTER_CUBIC) elif current_dpi 600: sigma 0.8 * (current_dpi / 300) img cv2.GaussianBlur(img, (0,0), sigmaXsigma) scale 300 / current_dpi img cv2.resize(img, None, fxscale, fyscale, interpolationcv2.INTER_LINEAR) return img5.2 笔迹类型迁移学习仅需50张样本的领域适配当目标场景为铅笔稿而非训练集的中性笔时全量微调成本高。项目提供轻量级适配方案冻结主干BiSeNetV2所有层requires_gradFalse替换头部将原mask_head替换为3层Conv32→16→1使用PSNRLoss.py中定义的PSNR-aware loss小批量训练batch_size4,epochs15,lr1e-5# 微调命令基于ckpt_convert.py生成的base模型 python train.py \ --pretrained_ckpt ckpt/best.pth \ --freeze_backbone True \ --head_only True \ --loss psnr \ --lr 1e-5 \ --epochs 15该方案在某省中考阅卷系统实测使用50张铅笔手写样本微调后擦除准确率从72.3%提升至89.6%耗时12分钟RTX 4090。5.3 教务系统集成要点API响应时间与错误码设计面向教育SaaS平台部署时需关注两点超时控制单张A4试卷2480×3508处理时间必须≤3.5秒含IO否则触发前端重试机制错误码分级400输入非图像文件检测magic bytes413图像尺寸10MB拒绝超大扫描件503GPU显存不足监控nvidia-smi显存占用95%时返回# Flask API核心逻辑predict.py app.route(/erase, methods[POST]) def erase_handwriting(): try: img_file request.files[image] img cv2.imdecode(np.frombuffer(img_file.read(), np.uint8), cv2.IMREAD_COLOR) if img is None: return jsonify({error: Invalid image format}), 400 # 显存健康检查 gpu_mem get_gpu_memory_usage() # 自定义函数 if gpu_mem 0.95: return jsonify({error: Server busy}), 503 result predict_image(model, img) # 调用前述predict_image return send_file(result, mimetypeimage/png) except Exception as e: app.logger.error(fPrediction failed: {str(e)}) return jsonify({error: Internal error}), 500实际部署中将get_gpu_memory_usage()改为查询nvidia-ml-py3库的nvmlDeviceGetMemoryInfo比nvidia-smi命令调用快17ms对高频请求至关重要。本文还有配套的精品资源点击获取
返回列表