ARTICLE DETAIL

资讯详情

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

PaddleOCR 中 CAN 手写数学公式识别算法:Counting-Aware Network 原理、配置与训练部署实战

PaddleOCR 中 CAN 手写数学公式识别算法:Counting-Aware Network 原理、配置与训练部署实战 PaddleOCR 中 CAN 手写数学公式识别算法Counting-Aware Network 原理、配置与训练部署实战【免费下载链接】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导读CANCounting-Aware Network计数感知网络是 PaddleOCR 中用于**手写数学公式识别HMER**的专用识别算法其核心思想是在序列解码过程中显式加入符号计数监督信号从而缓解传统注意力解码器在复杂手写公式上容易漏识别、错识别符号的问题。本文以 docs/version2.x/algorithm/formula_recognition/algorithm_rec_can.en.md 为骨架结合仓库中的完整配置文件与源码实现讲解 CAN 的算法原理、网络结构、数据与标签处理、训练/评估/预测全流程命令以及推理模型导出与 Python 端部署方法并附上关键参数逐一说明与 FAQ 排查建议。1. 算法背景与论文来源CAN 算法出自 ECCV 2022 论文When Counting Meets HMER: Counting-Aware Network for Handwritten Mathematical Expression RecognitionarXiv:2207.11463作者为 Bohan Li、Ye Yuan、Dingkang Liang 等。论文提出的核心观点是HMER 任务中一个常见的失败模式是注意力解码器遗漏公式中重复出现或密集分布的符号而统计每个符号出现次数这一相对简单的子任务可以为解码过程提供强大的辅助监督。PaddleOCR 在ppocr模块中完整复现了该算法使用CROHME 手写数学公式识别数据集进行训练与评估。当前仓库配置下在 CROHME 测试集上的复现效果表达式识别率 ExpRate如下表所示模型Backbone配置文件ExpRate预训练模型CANDenseNetrec_d28_can.yml51.72%rec_d28_can_train.tar训练产物见文档下载链接说明上表 ExpRate 数值取自仓库关联文档的复现记录对应 DenseNet 骨干网络 CANHead 的完整配置可作为复现实验的验收基准。2. 网络结构与源码级原理CAN 在 PaddleOCR 中被组织为标准识别模型三段式结构Transform空 BackboneDenseNet HeadCANHead并配套CANLoss损失函数、CANLabelEncode标签编码、CANLabelDecode结果解码与CANMetric评估指标。2.1 整体结构编码器 双分支解码器从 rec_can_head.py 的CANHead实现可以看到网络在骨干网络提取的 CNN 特征之上并行构建两条通路Counting Decoder计数分支由两个不同卷积核kernel_size3与kernel_size5的CountingDecoder构成输出每个符号类别的出现次数预测多尺度计数模块 MSCMAttention Decoder注意力解码分支AttDecoder在每一步解码时将计数预测结果通过counting_context_weight线性映射后与注意力上下文、隐状态、词嵌入一起融合指导符号序列的逐步生成。CANHead.forward的关键代码rec_can_head.py清晰地展示了这一融合流程def forward(self, inputs, targetsNone): cnn_features, images_mask, labels inputs counting_mask images_mask[:, :, ::self.ratio, ::self.ratio] counting_preds1, _ self.counting_decoder1(cnn_features, counting_mask) counting_preds2, _ self.counting_decoder2(cnn_features, counting_mask) counting_preds (counting_preds1 counting_preds2) / 2 word_probs self.decoder(cnn_features, labels, counting_preds, images_mask) return word_probs, counting_preds, counting_preds1, counting_preds2即计数预测取两个尺度分支的平均值再送入注意力解码器最终输出逐位置的符号概率分布。2.2 Counting Decoder如何数数CountingDecoderrec_can_head.py的实现包含三个子模块trans_layer1x1/3x3卷积 BatchNorm将特征通道压缩到 512ChannelAtt通道注意力模块AdaptiveAvgPool2D → FC → ReLU → FC → Sigmoid自适应地强调与计数相关的通道pred_layer1x1卷积输出 111 个通道对应 111 类符号经 Sigmoid 得到每个位置的符号概率图。最终在空间维度上求和paddle.sum(x, axis-1)得到每个符号的计数预测向量。这一预测与真实计数标签计算 SmoothL1 损失作为辅助监督信号。2.3 Attention Decoder带覆盖惩罚的注意力AttDecoderrec_can_head.py是一个基于 GRU 的注意力解码器具有以下特点位置编码使用PositionEmbeddingSine正弦位置编码num_pos_feats256注入编码器特征弥补纯卷积特征缺乏序列位置信息的缺陷覆盖惩罚CoverageAttention模块维护历史注意力权重累计和alpha_sum通过attention_conv11x11 卷积将覆盖历史映射为注意力偏置避免解码器反复关注同一区域导致符号重复或遗漏计数上下文融合每步解码状态由当前隐状态 词嵌入 注意力上下文 计数上下文四部分加权求和可加 Dropout 防止过拟合后经线性层映射为 111 维符号概率训练/推理行为差异is_trainTrue时使用真实标签逐位指导Teacher Forcingis_trainFalse时使用上一步预测的 argmax 结果自回归生成且默认解码步数为 36。2.4 骨干网络 DenseNetBackbone 采用 DenseNetrec_densenet.py3 个 DenseBlock每个 16 层Bottleneck 结构growthRate24配合 2 个 Transition 层做下采样输入为单通道灰度图。其特征图在高度和宽度方向上各缩小 16 倍与Head.ratio: 16对应最终输出 684 通道特征与Head.in_channel: 684严格匹配。2.5 损失函数符号级交叉熵 计数 SmoothL1CANLoss 由两部分组成word_average_loss符号序列的交叉熵损失默认不使用 label mask直接对全部位置求均值counting_loss三个计数预测两个分支 平均与真实计数标签的 SmoothL1 损失之和。真实计数标签由gen_counting_label生成统计每个符号在标签序列中出现的次数并忽略 0、1SOS/EOS 等特殊符与 107~110 号填充符rec_can_loss.py。总损失为二者直接相加。2.6 标签编码与解码LaTeX 符号字典训练阶段由CANLabelEncodelabel_ops.py完成将空格分隔的 LaTeX 符号序列如\frac { 1 } { 2 }逐符号映射为字典索引并追加结束符end_str预测阶段由CANLabelDecoderec_postprocess.py完成对每个 batch 找到结束符位置截断序列将索引还原为 LaTeX 符号并以空格连接输出符号字典为 ppocr/utils/dict/latex_symbol_dict.txt共 111 行对应out_channel: 111与word_num: 111。2.7 评估指标ExpRate 与 WordRateCANMetricrec_metric.py同时统计两个指标word_rate符号正确率基于SequenceMatcher计算预测序列与真实序列的相似度exp_rate表达式正确率ExpRate整条公式所有符号完全正确才计为 1是公式识别任务的黄金标准指标。配置中Metric.main_indicator: exp_rate表示以 ExpRate 作为早停与最优模型选择的依据。3. 环境准备与数据准备3.1 环境准备环境配置请参照 环境准备文档安装 PaddlePaddle、克隆 PaddleOCR 仓库并安装依赖。项目克隆方式参见 项目克隆文档。建议使用 GPU 环境进行训练与评估CPU 仅适合推理验证。3.2 CROHME 数据准备CAN 使用CROHME 手写数学公式识别数据集进行训练。数据目录组织方式与配置文件中的data_dir/label_file_list对应train_data/CROHME/ ├── training/ │ ├── images/ # 训练集手写公式图片 │ └── labels.txt # 训练集标签每行图片路径 空格分隔的LaTeX符号序列 └── evaluation/ ├── images/ # 评估集图片 └── labels.txt # 评估集标签标签文件每一行由图片路径 以空格分隔的 LaTeX 符号序列组成例如training/images/xxx.png \frac { 1 } { 2 } x ^ { 2 } y该格式与CANLabelEncode中label.strip().split()的解析逻辑label_ops.py完全对应。4. 模型训练 / 评估 / 预测PaddleOCR 将代码模块化训练不同的识别模型只需更换配置文件。CAN 的完整配置见 configs/rec/rec_d28_can.yml。4.1 配置要点速览配置项取值说明Global.epoch_num240总训练轮数Global.character_dict_pathppocr/utils/dict/latex_symbol_dict.txtLaTeX 符号字典111 类Global.max_text_length36标签最大长度与 Head 输出长度一致Global.use_space_charFalse空格不单独作为字符类别LaTeX 序列已含\space等显式符号Global.infer_imgdoc/datasets/crohme_demo/hme_00.jpg默认预测样例图OptimizerMomentum 0.9TwoStepCosine学习率 0.01warmup 1 epoch全局梯度裁剪 100.0weight_decay 1e-4Architecture.algorithmCAN算法标识用于路由到 CANHead 与 CANLossArchitecture.in_channels1单通道灰度输入BackboneDenseNetgrowthRate 24reduction 0.5bottleneck True16 层 DenseBlock x3Head.in_channel / out_channel684 / 111与 DenseNet 输出通道、字典大小严格对应Head.ratio16特征图相对输入的下采样倍数Head.attdecoderinput_size 256hidden_size 256attention_dim 512解码器核心超参is_train控制训练/推理模式LossCANLoss符号交叉熵 计数 SmoothL1PostProcessCANLabelDecodeLaTeX 序列解码MetricCANMetricmain_indicator: exp_rate以表达式识别率评估Train/EvalSimpleDataSet DyMaskCollator动态 mask 变长 batch 采集训练 batch_size 8、shuffle评估 batch_size 1其中两个关键工程细节值得注意DyMaskCollator手写公式图片与标签长度不一该 collate 函数负责动态构造 batch 内的有效区域 mask对应 Head 中的images_mask使变长序列可以批量训练灰度反转预处理GrayImageChannelFormat: inverse: True会将图片反色。CROHME 数据集多为深色背景浅色笔迹反转后统一为黑字白底与推理阶段--rec_image_inverse参数的含义保持一致。4.2 训练数据准备完成后即可启动训练# 单卡训练周期较长不推荐 python3 tools/train.py -c configs/rec/rec_d28_can.yml # 多卡训练通过 --gpus 指定 GPU 编号 python3 -m paddle.distributed.launch --gpus 0,1,2,3 tools/train.py -c configs/rec/rec_d28_can.yml训练产物默认保存在./output/rec/can/Global.save_model_dir每个 epoch 保存一次 checkpointsave_epoch_step: 1训练过程中每 1105 次迭代batch_size8 时约等于 1 个 epoch执行一次评估eval_batch_step: [0, 1105]并开启训练中指标计算cal_metric_during_train: True若需从断点继续训练可通过-o Global.checkpoints./output/rec/can/latest覆盖配置。4.3 评估# GPU 评估 python3 -m paddle.distributed.launch --gpus 0 tools/eval.py \ -c configs/rec/rec_d28_can.yml \ -o Global.pretrained_model./rec_d28_can_train/best_accuracy.pdparams评估使用训练保存的最优模型best_accuracy.pdparams输出 ExpRate 与 WordRate 指标。注意Global.pretrained_model应指向.pdparams权重文件评估/预测时需与训练保持同一份配置。4.4 预测单张图片# 预测用配置文件必须与训练保持一致 python3 tools/infer_rec.py \ -c configs/rec/rec_d28_can.yml \ -o Architecture.Head.attdecoder.is_trainFalse \ Global.infer_img./doc/datasets/crohme_demo/hme_00.jpg \ Global.pretrained_model./rec_d28_can_train/best_accuracy.pdparamstools/infer_rec.py是 PaddleOCR 面向单图/批量图片的识别脚本tools/infer_rec.py其内部调用Architecture.Head.attdecoder.is_trainFalse使解码器进入自回归生成模式解码步数固定 36。预测结果保存在Global.save_res_path默认./output/rec/predicts_can.txt。5. 推理模型导出与部署5.1 模型导出将训练得到的动态图权重转换为可部署的静态图推理模型python3 tools/export_model.py \ -c configs/rec/rec_d28_can.yml \ -o Global.pretrained_model./rec_d28_can_train/best_accuracy.pdparams \ Global.save_inference_dir./inference/rec_d28_can/ \ Architecture.Head.attdecoder.is_trainFalse # 模型默认输出最大长度为 36如需预测更长序列 # 请在导出时通过 Architecture.Head.max_text_length 指定合适的值如 # -o Architecture.Head.max_text_length72导出后将生成inference.pdmodel网络结构、inference.pdiparams权重与inference.yml配置三类文件存放于./inference/rec_d28_can/目录。注意注释中提到的是Architecture.Head.max_text_length但当前仓库配置中该参数实际位于Head顶层max_text_length: 36见 rec_d28_can.yml同时Global.max_text_length与Architecture.Head.max_text_length需保持一致。调整输出序列长度时请同步核对Global.max_text_length、Head.max_text_length与attdecoder相关配置并保证导出与推理使用同一参数集。5.2 Python 端推理使用通用识别预测脚本tools/infer/predict_rec.py通过参数指定算法与模型python3 tools/infer/predict_rec.py \ --image_dir./doc/datasets/crohme_demo/hme_00.jpg \ --rec_algorithmCAN \ --rec_batch_num1 \ --rec_model_dir./inference/rec_d28_can/ \ --rec_char_dict_path./ppocr/utils/dict/latex_symbol_dict.txt # 如需预测黑字白底白底黑字的图片请设置--rec_image_inverseFalse参数含义--rec_algorithmCAN指定使用 CAN 算法内部自动路由到CANHead与CANLabelDecode--rec_batch_num1CAN 解码器当前按固定序列长度输出建议 batch 数为 1 以确保稳定--rec_model_dir上一步导出的推理模型目录--rec_char_dict_pathLaTeX 符号字典必须与训练/导出时一致--rec_image_inverse是否对输入图片反色处理默认 True对应训练时GrayImageChannelFormat: inverse: True的数据增强方式若图片已是白底黑字应设为 False。5.3 其他部署形态按关联文档说明当前 CAN 算法在仓库中的部署支持情况如下C 推理不支持Serving 服务化部署不支持其他如移动端 Lite不支持。因此 CAN 当前主要以Python 端推理方式落地适用于科研复现与离线批量公式识别场景。6. FAQ 与常见问题排查预测结果为空或全部为结束符检查--rec_image_inverse设置是否与图片实际前景/背景色一致CROHME 训练数据经过反色处理若输入为白底黑字图片而未关闭反色特征分布会严重偏移。输出序列被截断结果明显短于公式实际长度模型默认输出长度为 36。对长公式请按 5.1 节说明在导出时增大max_text_length例如 72并同步更新Global.max_text_length与Head.max_text_length。训练 loss 不下降或计数分支异常确认标签文件为空格分隔的 LaTeX 符号序列格式且所有符号均存在于 latex_symbol_dict.txt111 类中不在字典中的符号会被CANLabelEncode静默跳过label_ops.py可能导致标签过短。显存不足调小Train.loader.batch_size_per_card默认 8或减少num_workersCAN 的注意力解码为逐步循环batch 过大会显著增加显存占用。复现指标低于文档 ExpRate 51.72%确认训练超参与 rec_d28_can.yml 完全一致含TwoStepCosine学习率、Momentum、240 epochs并使用与文档一致的 CROHME 训练/评估划分。7. 引用若在学术研究中使用 CAN 算法或本仓库实现请引用原论文misc{https://doi.org/10.48550/arxiv.2207.11463, doi {10.48550/ARXIV.2207.11463}, url {https://arxiv.org/abs/2207.11463}, author {Li, Bohan and Yuan, Ye and Liang, Dingkang and Liu, Xiao and Ji, Zhilong and Bai, Jinfeng and Liu, Wenyu and Bai, Xiang}, keywords {Computer Vision and Pattern Recognition (cs.CV), Artificial Intelligence (cs.AI), FOS: Computer and information sciences, FOS: Computer and information sciences}, title {When Counting Meets HMER: Counting-Aware Network for Handwritten Mathematical Expression Recognition}, publisher {arXiv}, year {2022}, copyright {arXiv.org perpetual, non-exclusive license} }结语CAN 是辅助监督 注意力解码思路在公式识别领域的代表性工作通过计数分支为注意力解码器注入符号级统计信息显著改善了手写公式中符号遗漏问题。在 PaddleOCR 中完整的 DenseNet 骨干、CANHead 双分支结构、CANLoss 组合损失与配套数据处理链路rec_d28_can.yml均已开源落地从数据准备到训练、评估、导出、Python 推理的整条流水线可直接运行。对公式识别、文档数学内容解析场景的开发者而言本文给出的配置参数对照、源码定位与部署命令即可作为开箱即用的实操手册。【免费下载链接】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),仅供参考
返回列表