PaddleOCR RobustScanner 文本识别算法解析:动态位置线索增强原理、配置详解与训练部署实战
【免费下载链接】PaddleOCR飞桨多语言OCR工具包(实用超轻量OCR系统,支持80+种语言识别,提供数据标注与合成工具,支持服务器、移动端、嵌入式及IoT设备端的训练与部署) Awesome multilingual OCR toolkits based on PaddlePaddle (practical ultra lightweight OCR system, support 80+ languages recognition, provide data annotation and synthesis tools, support training and deployment among server, mobile, embedded and IoT devices)项目地址: https://gitcode.com/paddlepaddle/PaddleOCR
RobustScanner(ECCV 2020)是 PaddleOCR 中面向不规则文本的注意力式识别算法,其核心思想是通过混合解码器动态增强位置线索,缓解纯序列注意力模型在规则文本上"注意力漂移"导致的识别退化问题。本文以 PaddleOCR 仓库中的 RobustScanner 复现文档为骨架,结合 rec_r31_robustscanner.yml 配置与 rec_robustscanner_head.py 源码实现,完整讲解算法原理、配置文件、数据流设计,以及从训练、评估、预测到 Python 推理部署的全流程实操,读者可直接照此在 PaddleOCR 中训练并部署 RobustScanner 识别模型。
1. 算法简介
1.1 论文信息
RobustScanner 出自 ECCV 2020 论文RobustScanner: Dynamically Enhancing Positional Clues for Robust Text Recognition(作者:Xiaoyu Yue, Zhanghui Kuang, Chenhao Lin, Hongbin Sun, Wayne Zhang)。
论文指出:序列注意力解码器(Sequence-to-Sequence Attention)在规则文本上容易因"注意力漂移"(attention drift)而失效,RobustScanner 通过显式增强**位置线索(Positional Clues)**来改善这一问题——同时维护一个"序列注意力解码器"(负责语义线索)和一个"位置注意力解码器"(负责位置线索),再用融合模块动态组合两者,从而在规则与不规则文本上均保持鲁棒。
1.2 复现效果
PaddleOCR 使用 MJSynth 和 SynthText 两个合成文字识别数据集训练,并在 IIIT、SVT、IC13、IC15、SVTP、CUTE 六个公开测试集上评估,复现效果如下:
| 模型 | 骨干网络 | 配置文件 | Acc | 预训练模型 |
|---|---|---|---|---|
| RobustScanner | ResNet31 | rec_r31_robustscanner.yml | 87.77% | 由官方文档提供训练权重下载(rec_r31_robustscanner.tar) |
注:除 MJSynth 与 SynthText 外,官方复现还额外使用了 SynthAdd 数据(百度网盘分享,提取码 627x)以及部分真实数据参与训练,具体数据细节可参考原论文。
2. 算法原理与源码实现
从源码结构看,RobustScanner 的完整解码流程位于 rec_robustscanner_head.py,由三部分构成:
2.1 编码器:ChannelReductionEncoder
RobustScannerHead首先通过ChannelReductionEncoder对骨干网络输出的高维特征做 1×1 卷积降维,将特征通道压缩为配置中的enc_outchannles: 128(代码第 703-705 行)。降维后的特征out_enc作为注意力机制的 Key(与 Query 计算相关性)。
2.2 混合解码器:hybrid decoder + position decoder + fusion
RobustScannerDecoder(代码第 524-681 行)内部维护两个并行的注意力解码器:
- SequenceAttentionDecoder(序列注意力解码器):以
<BOS>起始符号为输入,通过nn.Embedding+ 双层nn.LSTM生成 Query,与编码器特征做点积注意力(DotProductAttentionLayer),提取的是语义线索(glimpse)。 - PositionAttentionDecoder(位置注意力解码器):通过
PositionAwareLayer(对特征图按行做 LSTM 编码后接两个 3×3 卷积)得到"位置感知特征",再以显式的位置索引(word_positions,即 0..max_text_length-1 的位置序列)作为 Query,提取的是位置线索。
两个解码器输出的 glimpse 由RobustScannerFusionLayer(代码第 508-521 行)拼接后经线性层与 GLU 门控融合,最终送入分类层得到逐字符预测。训练时两者并行一次前向即可;测试时(forward_test)位置解码器一次性得到完整的位置 glimpse 序列,序列解码器则逐时间步自回归解码,每个 step 与对应位置 glimpse 融合后取 argmax,再回填到下一步输入。
2.3 注意力掩码:valid_ratio
DotProductAttentionLayer(代码第 97-123 行)支持按valid_ratio对注意力 logits 做掩码:对超出有效宽度的列填充-inf,经 softmax 后权重归零。该机制与预处理中按图片长宽比动态 padding 的设计配套,避免 padding 区域参与注意力计算。
3. 环境准备
请先参考 运行环境准备 配置 PaddleOCR 运行环境(安装 PaddlePaddle 与相关依赖),并参考 项目克隆 克隆项目代码。RobustScanner 的训练与推理均在 Python 端完成,无需额外编译自定义算子。
4. 配置文件详解
RobustScanner 的完整训练配置位于 rec_r31_robustscanner.yml。PaddleOCR 对代码进行了模块化,训练不同识别模型只需更换配置文件。下面按块逐项说明:
4.1 Global 全局配置
Global: use_gpu: true epoch_num: 5 log_smooth_window: 20 print_batch_step: 20 save_model_dir: ./output/rec/rec_r31_robustscanner/ save_epoch_step: 1 eval_batch_step: [0, 2000] # 每 2000 个 iter 评估一次 cal_metric_during_train: True pretrained_model: checkpoints: save_inference_dir: use_visualdl: False infer_img: doc/imgs_words_en/word_10.png character_dict_path: ppocr/utils/dict90.txt # 90 字符词典 max_text_length: &max_text_length 40 # 锚点,被 Head 与预处理复用 infer_mode: False use_space_char: False rm_symbol: True # 解码时移除符号、转小写 save_res_path: ./output/rec/predicts_robustscanner.txtcharacter_dict_path指向 dict90.txt,内含 90 个字符;经SARLabelEncode追加<UKN>、<BOS/EOS>、<PAD>三个特殊符后共 93 类(对应源码out_channels = 90 + unknown + start + padding,见 rec_robustscanner_head.py)。max_text_length: 40通过 YAML 锚点&max_text_length同时作用于 Head 与数据预处理。rm_symbol: True会让后处理SARLabelDecode用正则剔除英文字母、数字、中文以外的符号并将结果转小写(见 rec_postprocess.py),这是官方复现精度的关键配套设置。
4.2 Optimizer 优化器
Optimizer: name: Adam beta1: 0.9 beta2: 0.999 lr: name: Piecewise decay_epochs: [3, 4] values: [0.001, 0.0001, 0.00001] regularizer: name: 'L2' factor: 0采用 Adam 优化器与分段衰减学习率:epoch 3 前为 0.001,epoch 3-4 降为 0.0001,之后为 0.00001。
4.3 Architecture 网络结构
Architecture: model_type: rec algorithm: RobustScanner Transform: Backbone: name: ResNet31 init_type: KaimingNormal Head: name: RobustScannerHead enc_outchannles: 128 # 编码器降维通道数 hybrid_dec_rnn_layers: 2 # 序列注意力解码器 LSTM 层数 hybrid_dec_dropout: 0 position_dec_rnn_layers: 2 # 位置注意力解码器 LSTM 层数 start_idx: 91 # <BOS/EOS> 在 93 类词典中的索引 mask: True # 按 valid_ratio 掩码注意力 padding_idx: 92 # <PAD> 索引 encode_value: False # False 时注意力 value 使用原始特征而非编码器输出 max_text_length: *max_text_lengthstart_idx: 91与padding_idx: 92与SARLabelEncode中特殊字符的追加顺序严格对应(90 字符 +<UKN>(90) +<BOS/EOS>(91) +<PAD>(92),见 label_ops.py),改动词典时必须同步调整。
4.4 Loss 与 PostProcess
Loss: name: SARLoss PostProcess: name: SARLabelDecode Metric: name: RecMetric is_filter: True- SARLoss(rec_sar_loss.py):交叉熵损失,
ignore_index=92(即<PAD>不参与损失)。计算时丢弃模型输出的最后一位(与目标序列对齐),并丢弃目标序列的首位<BOS>。 - SARLabelDecode:解码时跳过
<PAD>并在遇到<EOS>时终止,支持rm_symbol清洗(rec_postprocess.py)。
4.5 Train / Eval 数据管线
训练与评估均使用LMDBDataSet,关键区别在RobustScannerRecResizeImg的image_shape: [3, 48, 48, 160](含义为:通道数 3、高 48、最小宽 48、最大宽 160):
Train: dataset: name: LMDBDataSet data_dir: ./train_data/data_lmdb_release/training/ transforms: - DecodeImage: img_mode: BGR channel_first: False - SARLabelEncode: - RobustScannerRecResizeImg: image_shape: [3, 48, 48, 160] width_downsample_ratio: 0.25 max_text_length: *max_text_length - KeepKeys: keep_keys: ['image', 'label', 'valid_ratio', 'word_positons'] loader: shuffle: True batch_size_per_card: 64 drop_last: True num_workers: 8 use_shared_memory: FalseEval 数据管线结构相同(data_dir指向 evaluation 目录,shuffle: False、drop_last: False、num_workers: 4),此处不再重复贴出。
RobustScannerRecResizeImg的实现位于 rec_img_aug.py:按原图宽高比将高度缩放到 48,width_downsample_ratio: 0.25意味着宽度必须是 4 的整数倍(width_divisor = int(1/0.25) = 4),并在 [48, 160] 范围内取整、裁剪;随后归一化并向右 padding 到固定宽 160(padding 值为 -1.0)。同时计算valid_ratio = min(1.0, resize_w / 160)供注意力掩码使用,并生成word_positons = [0, 1, ..., max_text_length-1]作为位置解码器的 Query。注意KeepKeys保留了word_positons(配置中拼写如此),这是 RobustScanner 区别于其他识别模型的关键输入。
5. 模型训练、评估与预测
完整的训练流程说明可参考 文本识别教程,核心命令如下。
5.1 训练
# 单卡训练(训练周期长,不建议) python3 tools/train.py -c configs/rec/rec_r31_robustscanner.yml # 多卡训练,通过 --gpus 参数指定卡号 python3 -m paddle.distributed.launch --gpus '0,1,2,3' tools/train.py -c configs/rec/rec_r31_robustscanner.yml训练过程每 2000 个 iter 评估一次,模型保存在./output/rec/rec_r31_robustscanner/下,best 权重文件名为best_accuracy。
5.2 评估
# GPU 评估,Global.pretrained_model 为待测权重 python3 -m paddle.distributed.launch --gpus '0' tools/eval.py -c configs/rec/rec_r31_robustscanner.yml -o Global.pretrained_model={path/to/weights}/best_accuracy5.3 预测
# 预测使用的配置文件必须与训练一致 python3 tools/infer_rec.py -c configs/rec/rec_r31_robustscanner.yml -o Global.pretrained_model={path/to/weights}/best_accuracy Global.infer_img=doc/imgs_words/en/word_1.png6. 推理与部署
6.1 Python 推理
第一步:导出 inference model。将训练保存的权重转换为推理模型:
python3 tools/export_model.py -c configs/rec/rec_r31_robustscanner.yml -o Global.pretrained_model={path/to/weights}/best_accuracy Global.save_inference_dir=./inference/rec_r31_robustscanner第二步:执行推理。predict_rec.py中rec_algorithm="RobustScanner"分支会自动装配SARLabelDecode并强制rm_symbol=True(见 predict_rec.py):
python3 tools/infer/predict_rec.py --image_dir="./doc/imgs_words/en/word_1.png" --rec_model_dir="./inference/rec_r31_robustscanner/" --rec_image_shape="3, 48, 48, 160" --rec_algorithm="RobustScanner" --rec_char_dict_path="ppocr/utils/dict90.txt" --use_space_char=False几个参数的要点:
--rec_image_shape="3, 48, 48, 160"是四维写法(通道、高、最小宽、最大宽),与训练配置中的image_shape保持一致;默认的--rec_image_shape是"3, 48, 320"(见 utility.py),RobustScanner 推理时必须显式覆盖为四维值,否则预处理宽度区间会与训练不一致。--rec_char_dict_path必须指向dict90.txt,与训练时的 93 类词典对齐。--use_space_char=False:RobustScanner 的 90 字符词典不含空格字符,需关闭空格字符。
6.2 C++ 推理
暂不支持。原因是 C++ 侧的预处理/后处理尚未覆盖 RobustScanner(SARLabelEncode、RobustScannerRecResizeImg等目前仅在 Python 数据管线中实现)。
6.3 Serving 服务化部署
暂不支持。
6.4 更多推理部署(Paddle Lite / ONNX 等)
暂不支持。
7. FAQ
Q:为什么
rec_image_shape是四个数字3, 48, 48, 160?RobustScanner 采用动态宽度预处理(高度固定 48,宽度按比例在 [48, 160] 内调整并 padding 到 160),因此配置中同时给出最小宽与最大宽。后两个数字分别对应imgW_min与imgW_max,见 rec_img_aug.py。Q:为什么
start_idx是 91、padding_idx是 92?因为词典为 90 字符 +<UKN>(90) +<BOS/EOS>(91) +<PAD>(92),这些索引由SARLabelEncode.add_special_char依次追加生成,Head 配置必须与其严格一致。Q:
encode_value=False有什么影响?此时注意力层的value使用骨干网络原始特征(dim_input),而非编码器降维后的特征(dim_model),分类层输入维度也随之变为dim_input;该开关在 rec_robustscanner_head.py 中体现。Q:训练时数据中
word_positons的作用是什么?它是位置注意力解码器的位置索引 Query,训练与推理时均由RobustScannerRecResizeImg生成(np.arange(0, max_text_length)),是位置线索的核心载体。Q:能否直接在 TIPC 流程中跑通 RobustScanner?仓库在 test_tipc/configs/rec_r31_robustscanner/ 中提供了对应的 TIPC 配置,可用于训练与推理的自动化验证。
引用
如使用 RobustScanner 算法,请引用:
@article{2020RobustScanner, title={RobustScanner: Dynamically Enhancing Positional Clues for Robust Text Recognition}, author={Xiaoyu Yue and Zhanghui Kuang and Chenhao Lin and Hongbin Sun and Wayne Zhang}, journal={ECCV2020}, year={2020}, }【免费下载链接】PaddleOCR飞桨多语言OCR工具包(实用超轻量OCR系统,支持80+种语言识别,提供数据标注与合成工具,支持服务器、移动端、嵌入式及IoT设备端的训练与部署) Awesome multilingual OCR toolkits based on PaddlePaddle (practical ultra lightweight OCR system, support 80+ languages recognition, provide data annotation and synthesis tools, support training and deployment among server, mobile, embedded and IoT devices)项目地址: https://gitcode.com/paddlepaddle/PaddleOCR
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考