PaddleOCR Text Gestalt 文本超分辨率实战:基于 Stroke-Aware 的 TSRN 模型训练、评估与推理部署
【免费下载链接】PaddleOCRTurn any PDF or image document into structured data for your AI. A powerful, lightweight OCR toolkit that bridges the gap between images/PDFs and LLMs. Supports 100+ languages.项目地址: https://gitcode.com/GitHub_Trending/pa/PaddleOCR
Text Gestalt(文本完形)是 PaddleOCR 提供的场景文本图像超分辨率(Scene Text Image Super-Resolution, STISR)算法,专门用于将模糊、低分辨率(LR)的文字图像重建为清晰的高分辨率(HR)图像,从而提升下游文字识别(OCR)的准确率。本文以 PaddleOCR 仓库中的 算法文档 为骨架,结合配置、源码与推理工具,完整讲解 Text Gestalt 的原理、数据准备、训练、评估、预测以及模型导出与部署的完整链路。读完本文,你将能够基于 PaddleOCR 独立复现 Text Gestalt 超分辨率模型,并掌握将训练权重转换为可部署推理模型的全部操作。
1. Text Gestalt 算法简介
Text Gestalt(论文Text Gestalt: Stroke-Aware Scene Text Image Super-Resolution,Chen, Jingye 等,发表于 AAAI 2022)是一类笔画感知(Stroke-Aware)的场景文本图像超分辨率算法。其核心思想是:普通图像超分只关注像素重建,而文本图像的超分必须同时关注文字笔画的结构完整性,才能在放大分辨率的同时保持文字可读性。
Paper: Text Gestalt: Stroke-Aware Scene Text Image Super-Resolution Chen, Jingye and Yu, Haiyang and Ma, Jianqi and Li, Bin and Xue, Xiangyang AAAI, 2022
在 PaddleOCR 中,Text Gestalt 的实现与 FudanOCR(复旦大学开源仓库)的text-gestalt分支一脉相承:ppocr/modeling/transforms/tsrn.py顶部明确标注This code is refer from: https://github.com/FudanVI/FudanOCR/blob/main/text-gestalt/model/tsrn.py,损失函数 stroke_focus_loss.py 同样源自 FudanOCR 的text-gestalt/loss/stroke_focus_loss.py,因此数据准备可完全参照 FudanOCR 的 TextZoom 数据下载说明。
1.1 TextZoom 测试集效果
参照 FudanOCR 的数据下载说明,Text Gestalt 超分算法在 TextZoom 测试集上的效果如下表所示:
| Model | Backbone | config | Acc | Download link |
|---|---|---|---|---|
| Text Gestalt | tsrn | configs/sr/sr_tsrn_transformer_strock.yml | 19.28 (PSNR+SSIM) | 0.6560 (识别准确率) |
其中:
- Backbone为
tsrn,即 TSRN(Text Super-Resolution Network)网络; - Acc(0.6560)表示超分后的图像送入识别模型得到的文字识别准确率;
- 19.28为 PSNR 与 SSIM 之和(
SRMetric中以all = psnr_avg + ssim_avg作为主指标,见下文源码分析); - 预训练权重可从官方 bcebos 链接
sr_tsrn_transformer_strock_train.tar下载(文档中给出的下载地址为https://paddleocr.bj.bcebos.com/sr_tsrn_transformer_strock_train.tar)。
1.2 从源码看 Text Gestalt 的网络构成
TSRN 网络定义于 ppocr/modeling/transforms/tsrn.py 的TSRN类,其结构可以拆解为以下三大部分:
SR 重建主干(Recurrent Residual Block + Upsample):
- 首个卷积层将输入从 3 通道映射到
2 * hidden_units(默认hidden_units=32)通道; - 中间堆叠
srb_nums=5个RecurrentResidualBlock(循环残差块,内部由Conv2D + BatchNorm2D + Mish + GruBlock组成); - 最后通过
UpsampleBLock(Conv2D + PixelShuffle + Mish)完成上采样,上采样倍数由scale_factor=2决定(upsample_block_num = int(math.log(scale_factor, 2))); - 输出经
paddle.tanh归一化得到sr_img。
- 首个卷积层将输入从 3 通道映射到
STN 空间变换(可选):当配置中
STN: True时,TSRN 会先通过STN_model预测控制点,再经TPSSpatialTransformer(TPS 薄板样条变换)对输入进行形变矫正,缓解场景文本的透视/弯曲问题。源码中tps_inputsize = [height // scale_factor, width // scale_factor],即 STN 在低分辨率尺度上工作。Transformer 识别分支(训练时冻结):TSRN 内嵌一个
r34_transformer = Transformer()(定义于 ppocr/modeling/heads/sr_rensnet_transformer.py),且构造后所有参数trainable = False。训练时它对超分图sr_img和高清图hr_img分别做识别,产出sr_pred/hr_pred以及word_attention_map_pred/word_attention_map_gt两组注意力图,用于计算笔画聚焦损失。推理时该分支不参与计算(if self.training:分支跳过)。
TSRN.forward的输出字典在训练阶段包含sr_img、hr_img、hr_pred、sr_pred、word_attention_map_gt、word_attention_map_pred等键;在推理阶段(infer_mode=True)仅输出lr_img与sr_img,这也与 tools/infer_sr.py 中preds["sr_img"]、preds["lr_img"]的取值方式一一对应。
1.3 StrokeFocusLoss 笔画聚焦损失
训练使用的损失函数是 StrokeFocusLoss,其前向计算逻辑为:
mse_loss = self.mse_loss(sr_img, hr_img) attention_loss = paddle.nn.functional.l1_loss( word_attention_map_gt, word_attention_map_pred ) loss = (mse_loss + attention_loss * 50) * 100即总损失由两部分组成:
- 像素级 MSE 损失:约束超分图与高清图在像素层面一致;
- 笔画注意力 L1 损失(权重 50):约束超分图上的文字笔画注意力图与高清图一致,这正是 Text Gestalt"笔画感知"的核心体现;
- 整体再乘以 100 放大梯度尺度。
损失内部维护的english_stroke_dict将字符映射到笔画分解序列数字(见 2.3 节),与SRLabelEncode中的映射保持一致。
1.4 SRMetric 评估指标
评估指标SRMetric定义于 ppocr/metrics/sr_metric.py,它同时计算并累计 PSNR 与 SSIM:
calculate_psnr:基于 MSE 计算峰值信噪比,mse == 0时返回inf;calculate_ssim:使用内置的SSIM类计算结构相似度;get_metric返回{"psnr_avg": ..., "ssim_avg": ..., "all": psnr_avg + ssim_avg},其中all即配置文件中Metric.main_indicator: all对应的主指标——文档表格中 Text Gestalt 的19.28正是该all值。
2. 环境准备与数据准备
2.1 环境配置与项目克隆
请参照 环境准备 配置 PaddleOCR 运行环境(包括安装 PaddlePaddle 与 PaddleOCR 依赖),并参照 项目克隆 克隆当前仓库代码。
2.2 数据准备:TextZoom
Text Gestalt 训练/评估使用 TextZoom 数据集。参照 FudanOCRtext-gestalt分支的数据下载说明完成下载后,将数据整理为 PaddleOCR 期望的 LMDB 格式,并放置到配置指定的路径下:
- 训练数据目录:
./train_data/srdata/train(配置Train.dataset.data_dir) - 测试数据目录:
./train_data/srdata/test(配置Eval.dataset.data_dir)
数据集中的每一条样本应包含低分辨率图image_lr、高清图image_hr以及对应的文字标签label,供 LMDBDataSetSR 读取。
2.3 字符笔画分解字典
配置中两处引用了./train_data/srdata/english_decomposition.txt(Global.character_dict_path 与Loss.character_dict_path)。该字典的格式为每行一个字符及其笔画分解序列,例如:
A 012345 B 023456 ...从 SRLabelEncode 的加载逻辑(character, sequence = line.split())与StrokeFocusLoss的读取方式(character, sequence = line.split())可以看到:字典将字符映射为笔画数字序列(0-9 数字),标签编码时把单词的每个字符替换为笔画序列并以"0"结尾,得到stroke_sequence,再按english_stroke_dict = "0123456789"映射为整数张量input_tensor与长度length。这一"字符 → 笔画序列"编码是笔画聚焦机制的数据基础。
3. 模型训练 / 评估 / 预测
PaddleOCR 采用模块化设计,训练不同模型只需修改配置文件。Text Gestalt 的完整配置见 configs/sr/sr_tsrn_transformer_strock.yml,其关键模块说明如下:
| 配置模块 | 值 | 说明 |
|---|---|---|
Global.model_type | sr | 模型类型为超分辨率 |
Architecture.algorithm | Gestalt | 算法标识为 Gestalt |
Architecture.Transform.name | TSRN | 网络主干 |
Architecture.Transform.STN | True | 启用 TPS 空间变换矫正 |
Loss.name | StrokeFocusLoss | 笔画聚焦损失 |
PostProcess.name | None | 无后处理 |
Metric.name | SRMetric | PSNR + SSIM 指标 |
Metric.main_indicator | all | 主指标为 PSNR 与 SSIM 之和 |
Train.dataset.name | LMDBDataSetSR | 训练数据读取器 |
数据增强管道(Train.dataset.transforms与Eval.dataset.transforms)包含三个算子:
SRResize:imgH: 32, imgW: 128, down_sample_scale: 2。从 operators.py 中的 SRResize 看,它把高清图缩放到(128, 32),把低清图缩放到(imgW // down_sample_scale, imgH // down_sample_scale) = (64, 16),即内部生成 2 倍降采样后的 LR 图;infer_mode=True时只保留img_lr。SRLabelEncode:基于笔画分解字典做标签编码,产出length与input_tensor。KeepKeys:按['img_lr', 'img_hr', 'length', 'input_tensor', 'label']顺序组织 DataLoader 返回值。
训练超参数方面:Global.epoch_num: 500、save_epoch_step: 3、eval_batch_step: [0, 1000](每 1000 个 iteration 评估一次);优化器为 Adam(beta1: 0.5, beta2: 0.999, clip_norm: 0.25),学习率0.0001;Train.loader.batch_size_per_card: 16。
3.1 模型训练
数据准备完成后即可开始训练:
# 单卡训练(训练周期较长,不推荐) python3 tools/train.py -c configs/sr/sr_tsrn_transformer_strock.yml # 多卡训练,通过 --gpus 指定使用的 GPU 编号 python3 -m paddle.distributed.launch --gpus '0,1,2,3' tools/train.py -c configs/sr/sr_tsrn_transformer_strock.yml3.2 模型评估
# GPU 评估 python3 -m paddle.distributed.launch --gpus '0' tools/eval.py -c configs/sr/sr_tsrn_transformer_strock.yml -o Global.pretrained_model={path/to/weights}/best_accuracy其中Global.pretrained_model指向训练过程中保存的最佳权重(save_model_dir下的best_accuracy),评估输出即SRMetric的psnr_avg、ssim_avg与all。
3.3 模型预测
预测使用专用的 tools/infer_sr.py:
# 预测所用配置文件必须与训练一致 python3 tools/infer_sr.py -c configs/sr/sr_tsrn_transformer_strock.yml -o Global.pretrained_model={path/to/weights}/best_accuracy Global.infer_img=doc/imgs_words_en/word_52.pnginfer_sr.py的推理流程(对应 源码 L50-L94)为:
- 将
Architecture.Transform.infer_mode置为True,构建模型并加载权重; - 复用
Eval.dataset.transforms构建预处理算子,跳过SRLabelEncode,并令SRResize.infer_mode=True、KeepKeys只保留img_lr; - 逐张读取
Global.infer_img指定的图片,前向得到sr_img与lr_img; - 将结果乘以 255 并转置为 HWC 格式,保存为
infer_result/sr_{原文件名}(保存目录由Global.save_visual控制,默认infer_result/)。
输入图doc/imgs_words_en/word_52.png超分结果示例:
4. 推理与部署
4.1 Python 推理
首先需要把训练过程保存的模型导出为推理模型(也可直接使用官方sr_tsrn_transformer_strock_train.tar权重):
python3 tools/export_model.py -c configs/sr/sr_tsrn_transformer_strock.yml -o Global.pretrained_model={path/to/weights}/best_accuracy Global.save_inference_dir=./inference/sr_out导出完成后,使用 tools/infer/predict_sr.py 进行 Text Gestalt 超分推理:
python3 tools/infer/predict_sr.py --sr_model_dir=./inference/sr_out --image_dir=doc/imgs_words_en/word_52.png --sr_image_shape=3,32,128参数说明:
--sr_model_dir:导出后的推理模型目录;--image_dir:待超分的图像路径或目录;--sr_image_shape:通道数,高,宽,默认3,32,128,即期望的超分输出尺寸;--sr_batch_num:推理批大小(见 predict_sr.py 的 TextSR 类)。
从 TextSR.resize_norm_img 可以看到,Python 推理时输入图先被缩放到(imgW // 2, imgH // 2) = (64, 16)作为 LR 输入(与训练时down_sample_scale: 2的预处理保持一致),经create_predictor创建 Paddle Inference 预测器后前向得到超分结果。执行上述命令后,word_52.png的超分结果如下:
4.2 C++ 推理
暂不支持。
4.3 Serving 服务化部署
暂不支持。
4.4 更多部署方式
暂不支持。
5. 常见问题(FAQ)
- 训练时显存不足怎么办?可调低
Train.loader.batch_size_per_card(默认 16)或num_workers,并配合Global.print_batch_step观察训练进度。 - 推理输出与预期尺寸不符?请检查
--sr_image_shape是否与训练配置中的SRResize.imgH/imgW一致(默认3,32,128),且推理配置文件必须与训练一致。 - 识别准确率提升不明显?Text Gestalt 是针对场景文本图像(如 TextZoom 数据)设计的,对极端模糊或畸变文本,建议配合 STN(
STN: True)与充足的低-高清成对数据训练。
引用
如果 Text Gestalt 对你的研究或工作有帮助,请引用原论文:
@inproceedings{chen2022text, title={Text gestalt: Stroke-aware scene text image super-resolution}, author={Chen, Jingye and Yu, Haiyang and Ma, Jianqi and Li, Bin and Xue, Xiangyang}, booktitle={Proceedings of the AAAI Conference on Artificial Intelligence}, volume={36}, number={1}, pages={285--293}, year={2022} }延伸阅读
- 同目录下的另一超分算法:Text Telescope 超分辨率算法文档;
- 文本识别训练通用流程:Text Recognition Tutorial;
- 更多训练/推理脚本参见 tools/train.py、tools/eval.py 与 tools/export_model.py。
【免费下载链接】PaddleOCRTurn any PDF or image document into structured data for your AI. A powerful, lightweight OCR toolkit that bridges the gap between images/PDFs and LLMs. Supports 100+ languages.项目地址: https://gitcode.com/GitHub_Trending/pa/PaddleOCR
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考