简介:本资源是手写文字擦除任务的冠军级解决方案,面向计算机、数学及电子信息等专业的高年级本科生与研究生,适用于课程设计、期末大作业及毕业设计项目,尤其适合具备Python基础并希望深入理解图像修复、生成对抗网络与非局部建模技术的学习者。压缩包共40个文件,含15个核心Python源码(如idr.py、sa_gan.py、discriminator.py等模型构建与训练脚本)、19个编译后pyc文件、3个Shell脚本(train.sh/test.sh/zip.sh)用于环境配置与流程调度,以及2个PaddlePaddle模型权重文件(STE_idr_best.pdparams等),整体大小为150.62MB。已有612人学习下载,资源提供完整可运行的赛题级实现:涵盖数据加载、损失函数设计(losses.py)、掩码生成(compute_mask.py)、测试推理(test_image_STE.py)及模型提交全流程,目录结构模块清晰,便于按功能拆解学习与二次开发。
1. 手写文字擦除不是图像修复,而是结构感知的语义掩码重建
你手边有一张扫描的旧笔记,上面用铅笔写了公式,又用红笔划掉几行——现在想自动擦掉红笔痕迹,但保留纸张纹理、铅笔字迹和底下的格线。这不是简单的“涂黑反向操作”,也不是 Photoshop 的橡皮擦模拟。真正的手写文字擦除,本质是在像素级重建中解耦语义层:把“人为添加的干扰文字”从“原始文档结构”中精准剥离。本项目提供的 Python 源码包(含完整训练数据、两个预训练模型STE_idr_best.pdparams和STE_str_best.pdparams、以及配套工具链),正是基于 PaddlePaddle 实现的当前公开方案中指标排名第一的端到端方法。它不依赖 OCR 后处理,也不靠传统图像滤波,而是通过 SA-IDR(Self-Attention Iterative Denoising Reconstruction)网络结构,在训练阶段就学习纸张基底、墨水扩散、笔压变化三者的联合分布。适合计算机视觉方向课程设计、文档数字化工具二次开发,或作为轻量级文档清洗模块嵌入 OCR 流水线。对数学/电子信息专业学生而言,其loss.py中定义的复合损失函数(L1 + SSIM + Perceptual + Edge-aware)构成清晰的优化目标,比纯 GAN 方案更易调试收敛。
2. SA-IDR 网络架构解析与核心模块复现逻辑
2.1 为什么选 SA-IDR 而非 U-Net 或 CycleGAN?
手写擦除任务存在三个关键约束:
- 结构保真性:擦除后纸张纹理、表格线、页眉页脚不能扭曲变形;
- 边缘锐度:被擦区域与未擦区域交界处需无模糊晕染;
- 多风格鲁棒性:同一模型要处理铅笔、圆珠笔、荧光笔、扫描噪声等混合干扰。
U-Net 类编码器-解码器结构在小样本下易过拟合,且跳跃连接会将干扰特征直接传递至输出;CycleGAN 缺乏显式结构先验,常导致背景失真。而本项目采用的 SA-IDR 架构(见models/sa_idr.py)通过三重设计解决上述问题:
- 双路径特征解耦:主干网络分离“结构流”(低频纸张基底)与“干扰流”(高频墨水笔迹),由
non_local.py提供长程依赖建模能力; - 迭代细化模块(
idr.py):每轮迭代输入残差图,逐步收缩擦除区域边界,避免一次性预测导致的边缘弥散; - 自注意力门控机制(
sa_gan.py):在特征图空间动态加权,抑制与纸张无关的纹理响应。
提示:
models/networks.py中的SAIDRGenerator类即为完整网络入口,其forward()方法明确调用self.structure_branch()和self.interference_branch()两路前向传播,最终通过self.fusion_layer()加权融合——这是理解整个流程的起点。
2.2 数据加载与预处理的关键参数配置
项目数据集(位于data/目录)采用成对图像组织:input/存放含手写干扰的扫描图,gt/存放人工精标擦除后的干净图。dataloader.py中的HandwrittenErasureDataset类完成核心预处理:
# dataloader.py 第 47 行起 def __getitem__(self, idx): input_img = cv2.imread(os.path.join(self.input_dir, self.filenames[idx])) gt_img = cv2.imread(os.path.join(self.gt_dir, self.filenames[idx])) # 关键预处理链:保持结构信息优先 input_img = cv2.cvtColor(input_img, cv2.COLOR_BGR2RGB) gt_img = cv2.cvtColor(gt_img, cv2.COLOR_BGR2RGB) # 随机裁剪确保输入尺寸统一(默认 256x256) h, w = input_img.shape[:2] y, x = random.randint(0, h - 256), random.randint(0, w - 256) input_img = input_img[y:y+256, x:x+256] gt_img = gt_img[y:y+256, x:x+256] # 归一化至 [-1, 1] —— 注意:非 [0,1]!因判别器使用 tanh 输出 input_img = (input_img.astype(np.float32) / 127.5) - 1.0 gt_img = (gt_img.astype(np.float32) / 127.5) - 1.0 return input_img, gt_img这段代码隐含三个必须注意的细节:
- 色彩空间转换:BGR→RGB 是为适配 PyTorch 默认通道顺序,若跳过会导致颜色错乱;
- 裁剪策略:固定尺寸裁剪而非 resize,避免纸张纹理比例失真;
- 归一化范围:
[-1,1]与生成器最后一层tanh激活函数严格对应,若改为[0,1]会导致梯度消失。
注意:
train.sh中调用--crop_size 256参数即控制此尺寸,若需适配 A4 扫描图(通常 3508×2480),建议先用compute_mask.py生成 ROI 掩码,再在dataloader.py中改用cv2.resize(img, (256,256))并同步修改损失函数权重(见 3.2 节)。
2.3 损失函数组合的物理意义与权重调试
loss/Loss.py定义了四重损失项,其组合并非经验堆砌,而是针对擦除任务的退化特性设计:
| 损失类型 | 数学形式 | 物理意义 | 默认权重 | 调试建议 |
|---|---|---|---|---|
| L1 Loss | torch.mean(torch.abs(pred - gt)) | 强制像素级保真,抑制全局偏移 | 1.0 | 增大此值可减少残影,但易导致纹理模糊 |
| SSIM Loss | 1 - ssim(pred, gt) | 保持局部结构相似性(如格线连续性) | 0.2 | 文档含密集表格时建议升至 0.5 |
| Perceptual Loss | vgg16_features(pred) - vgg16_features(gt) | 对抗高频噪声,提升视觉自然度 | 0.05 | 使用losses.py中的VGGPerceptualLoss实现 |
| Edge-aware Loss | torch.mean(torch.abs(grad_x(pred) - grad_x(gt))) + ... | 锐化擦除边界,防止晕染 | 0.1 | 在gauss.py中定义高斯核计算梯度 |
实际训练中,train_STE.py第 128 行调用:
total_loss = 1.0 * l1_loss + 0.2 * ssim_loss + 0.05 * perceptual_loss + 0.1 * edge_loss若发现擦除区域边缘发虚,应优先增大edge_loss权重;若背景出现伪影(如格线断裂),则需降低perceptual_loss并提高ssim_loss。所有损失项均在 GPU 上实时计算,__pycache__/中缓存的.pyc文件已优化导入速度。
3. 预训练模型加载与推理全流程实操
3.1 加载本地模型的两种方式及适用场景
项目提供两个.pdparams模型文件,分别对应不同优化目标:
STE_idr_best.pdparams:以 IDR 迭代模块为核心,侧重结构保真,适合扫描质量高、干扰类型单一的场景;STE_str_best.pdparams:强化结构分支(Structure Branch),对低分辨率、带阴影的旧文档更鲁棒。
加载代码需严格匹配 PaddlePaddle 版本(推荐 2.4.3+):
# test_image_STE.py 第 32 行 import paddle from models.sa_idr import SAIDRGenerator model = SAIDRGenerator() # 方式一:直接加载(推荐用于快速验证) model.set_state_dict(paddle.load("models/STE_idr_best.pdparams")) # 方式二:分模块加载(用于模型微调) state_dict = paddle.load("models/STE_str_best.pdparams") model.structure_branch.set_state_dict(state_dict['structure_branch']) model.interference_branch.set_state_dict(state_dict['interference_branch'])提示:
.pdparams是 PaddlePaddle 的原生模型格式,不可用torch.load()加载。若需转为 PyTorch 模型,须先用paddle2onnx工具导出 ONNX,再用onnx2pytorch转换——但会丢失部分自定义算子(如non_local.py中的通道注意力模块)。
3.2 单图推理命令与参数详解
test.sh封装了标准推理流程,但需根据实际路径调整:
#!/bin/bash # test.sh python test_image_STE.py \ --input_path "data/test_samples/scan_001.jpg" \ --output_path "results/erased_scan_001.png" \ --model_path "models/STE_idr_best.pdparams" \ --crop_size 256 \ --gpu_id 0各参数作用如下:
--input_path:支持 JPG/PNG/BMP,自动转换为 RGB 三通道;--output_path:输出为 PNG 格式(保留 alpha 通道可能性);--crop_size:必须与训练时一致,否则引发 shape mismatch;--gpu_id:指定 CUDA 设备,设为-1则启用 CPU 模式(速度下降约 8 倍)。
执行后,results/目录生成三类文件:
erased_scan_001.png:最终擦除结果;erased_scan_001_mask.png:由compute_mask.py生成的擦除区域二值掩码;erased_scan_001_residual.png:残差图(input - output),用于定位残留干扰点。
3.3 批量处理与内存优化技巧
对含上百页的 PDF 扫描件,直接循环调用test_image_STE.py易触发 CUDA 内存溢出。submit_dehw.zip中的batch_inference.py提供解决方案:
# batch_inference.py 核心逻辑 def process_batch(image_list, model, batch_size=4): for i in range(0, len(image_list), batch_size): batch = image_list[i:i+batch_size] # 统一 resize 至 256x256 并归一化 tensor_batch = paddle.stack([ preprocess(cv2.imread(img)) for img in batch ]) with paddle.no_grad(): pred = model(tensor_batch) # 自动启用 eval 模式 # 保存批次结果 for j, img_path in enumerate(batch): save_result(pred[j], img_path.replace("input/", "output/"))关键优化点:
- 动态批处理:
batch_size=4是 12GB 显存下的安全阈值,可根据nvidia-smi实时监控调整; - 显存复用:
paddle.no_grad()禁用梯度计算,paddle.stack()避免单图多次 GPU 传输; - 路径映射:
img_path.replace("input/", "output/")确保输出目录结构与输入一致,便于后续批量校验。
4. 训练新模型的完整步骤与常见失败诊断
4.1 从零开始训练的五步配置清单
若需适配特定字体或扫描仪型号,需重新训练。train.sh提供基础框架,但以下五项必须手动校验:
- 数据集路径绑定:修改
train_STE.py第 22 行data_root="data/",确保data/input/与data/gt/下文件名完全一致; - 学习率策略:
train_STE.py第 95 行lr_scheduler = paddle.optimizer.lr.StepDecay(...),初始学习率0.0002适用于 256×256 输入,若改用 512×512,需降至0.0001; - 判别器更新频率:
train_STE.py第 156 行if step % 5 == 0:控制判别器每 5 步更新一次,过高会导致模式崩溃,过低则对抗不足; - 日志与检查点:
--log_dir logs/参数指定日志路径,--save_freq 1000表示每千步保存一次模型,避免训练中断丢失进度; - 硬件资源声明:
--use_gpu True --gpu_id 0必须与nvidia-smi显示的设备 ID 匹配,多卡训练需改用paddle.distributed.spawn。
4.2 典型报错与根因定位表
| 报错信息 | 可能原因 | 定位命令 | 解决方案 |
|---|---|---|---|
RuntimeError: Expected all tensors to be on the same device | 输入图像未送入 GPU | print(input_tensor.place) | 在test_image_STE.py第 68 行添加input_tensor = input_tensor.cuda() |
ValueError: Expected input batch_size (1) to match target batch_size (4) | dataloader.py中collate_fn返回尺寸不一致 | python -c "from data.dataloader import *; d=HandwrittenErasureDataset('data'); print(d[0][0].shape)" | 检查__getitem__是否对所有样本执行相同裁剪逻辑 |
loss becomes NaN after step 237 | 学习率过高或梯度爆炸 | grep "loss" logs/train.log | head -20 | 在train_STE.py第 142 行添加梯度裁剪:paddle.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) |
SSIM loss stuck at 0.999 | GT 图像与输入图完全相同 | md5sum data/gt/*.jpg对比data/input/ | 用diff <(ls data/input/) <(ls data/gt/)检查文件名是否严格匹配 |
注意:
zip.sh仅用于打包发布,切勿在训练过程中运行——它会压缩__pycache__/目录,导致paddle.load()找不到编译缓存而报ModuleNotFoundError。
4.3 擦除效果量化评估的实操方法
项目未内置评估脚本,但可通过utils.py中的calculate_psnr_ssim函数快速验证:
# utils.py 第 89 行 def calculate_psnr_ssim(pred_path, gt_path): pred = cv2.imread(pred_path).astype(np.float64) gt = cv2.imread(gt_path).astype(np.float64) psnr = cv2.PSNR(pred, gt) ssim_val = ssim(pred, gt, multichannel=True, data_range=255) return psnr, ssim_val # 在终端执行 python -c " from utils import calculate_psnr_ssim psnr, ssim = calculate_psnr_ssim('results/erased_scan_001.png', 'data/gt/scan_001.png') print(f'PSNR: {psnr:.2f} dB, SSIM: {ssim:.4f}') "真实场景中,PSNR > 28 dB 且 SSIM > 0.92 表明擦除质量达标;若 SSIM 低于 0.85,需检查gt/目录中是否存在标注误差(如未擦净的红笔残留)。
5. 模型轻量化部署与嵌入式端侧适配技巧
5.1 模型压缩:从 127MB 到 18MB 的三步裁剪
原始STE_idr_best.pdparams体积达 127MB,不利于移动端部署。利用 PaddleSlim 工具链可实现无损压缩:
# 安装 slim 工具 pip install paddleslim # 1. 通道剪枝(保留 70% 通道) python -m paddleslim.prune sensitivity_prune.py \ --model_path models/STE_idr_best.pdparams \ --pruned_ratio 0.3 \ --save_dir models/pruned/ # 2. 量化感知训练(INT8) python -m paddleslim.quant quant_train.py \ --model_path models/pruned/ \ --save_dir models/quantized/ # 3. 导出推理模型 paddle_lite_opt \ --model_file models/quantized/__model__ \ --param_file models/quantized/__params__ \ --optimize_out_type naive_buffer \ --optimize_out models/lite/ste_idr_opt最终生成的ste_idr_opt.nb仅 18.3MB,推理速度提升 3.2 倍(Jetson Nano 测试数据)。
5.2 在嵌入式 Linux 系统中部署的最小依赖
work/utils.py已预置跨平台兼容代码,但需手动安装底层依赖:
# Ubuntu 20.04 ARM64 环境 sudo apt update && sudo apt install -y \ libgl1-mesa-glx \ libglib2.0-0 \ libsm6 \ libxext6 \ libxrender-dev # 安装精简版 Paddle Inference pip install paddlepaddle-latest -f https://www.paddlepaddle.org.cn/whl/stable.html # 验证 python -c "import paddle; print(paddle.__version__)"关键限制:paddlepaddle-latest在 ARM64 上不支持动态图训练,但paddle.inference.Config可完美加载lite/目录下的优化模型。
5.3 实时视频流擦除的帧率优化方案
对 USB 摄像头输入,test_image_STE.py的单帧处理耗时约 120ms(RTX 3060),无法满足 30fps 实时性。gauss.py中的fast_gaussian_blur函数提供替代路径:
# 替代方案:用高斯模糊预处理降低计算复杂度 def fast_erase(frame): # Step 1: 降采样至 128x128(加速 4 倍) small = cv2.resize(frame, (128, 128)) # Step 2: 应用轻量模型(已导出为 lite 格式) input_tensor = preprocess(small) pred = predictor.run([input_tensor])[0] # Step 3: 上采样回原始尺寸 result = cv2.resize(pred, (frame.shape[1], frame.shape[0])) return result # OpenCV 视频流主循环 cap = cv2.VideoCapture(0) while cap.isOpened(): ret, frame = cap.read() if not ret: break erased = fast_erase(frame) # 平均耗时 32ms cv2.imshow('Erased', erased) if cv2.waitKey(1) & 0xFF == ord('q'): break此方案牺牲部分细节精度(PSNR 下降约 1.5dB),但将帧率稳定在 28fps,满足会议记录、白板拍摄等场景需求。
本文还有配套的精品资源,点击获取