PaddleOCR 手写数学公式识别算法 CAN 实战指南:Counting-Aware Network 训练、评估与推理部署
【免费下载链接】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
手写数学公式识别(HMER)是 OCR 领域中极具挑战性的任务,其难点在于公式的二维空间结构、符号歧义与书写随意性。本指南以 PaddleOCR 仓库中的 CAN(Counting-Aware Network)算法文档为核心,完整讲解该算法在 PaddleOCR 中的训练、评估、预测与推理部署全流程,并结合 rec_d28_can.yml 配置与 rec_can_head.py 等源码,深入剖析其 Counting 模块与 Attention Decoder 的实现原理,帮助读者从"会跑命令"进阶到"理解算法"。
1. 算法简介
CAN(Counting-Aware Network)由 Bohan Li、Ye Yuan、Dingkang Liang、Xiao Liu、Zhilong Ji、Jinfeng Bai、Wenyu Liu、Xiang Bai 等人提出,论文《When Counting Meets HMER: Counting-Aware Network for Handwritten Mathematical Expression Recognition》发表于 ECCV 2022。其核心思想是:在传统的序列到序列(Seq2Seq)识别框架之外,额外引入一个"计数解码器(Counting Decoder)",显式地统计每个数学符号在图像中出现的次数,以此约束注意力机制,缓解手写公式中符号密集、空间错位导致的漏识别与错识别问题。
PaddleOCR 中 CAN 使用 CROHME 手写公式数据集训练,对应测试集上的精度如下:
| 模型 | 骨干网络 | 配置文件 | ExpRate | 下载链接 |
|---|---|---|---|---|
| CAN | DenseNet | rec_d28_can.yml | 51.72% | 训练模型 |
说明:ExpRate(Expression Recognition Rate)是公式级识别准确率,即整条公式的符号序列完全正确的比例,比字符级准确率更为严格。
2. 网络结构与源码实现
2.1 整体架构
CAN 在 PaddleOCR 中遵循"Backbone + Head"的模块化设计,由 rec_d28_can.yml 中的Architecture字段定义:
- Backbone:
DenseNet,配置growthRate: 24、reduction: 0.5、bottleneck: True、use_dropout: True、input_channel: 1,输入为单通道灰度图; - Head:
CANHead,in_channel: 684(DenseNet 输出的特征通道数)、out_channel: 111(符号类别数)、max_text_length: 36、ratio: 16(特征图相对原图的下采样倍数)。
2.2 Counting 模块:多尺度计数解码器
在 rec_can_head.py 中,CANHead内部构造了两个CountingDecoder,分别使用卷积核大小为 3 和 5 的trans_layer提取特征,并通过ChannelAtt通道注意力(自适应平均池化 + 两层全连接 + Sigmoid)进行通道加权,最后以 1×1 卷积加 Sigmoid 输出每个符号的"计数热力图",再按空间维度求和得到符号计数预测:
counting_preds1:kernel_size=3 的计数解码器输出;counting_preds2:kernel_size=5 的计数解码器输出;counting_preds = (counting_preds1 + counting_preds2) / 2:两者取平均作为最终计数向量。
多尺度卷积核分别关注局部与更广感受野的符号分布,增强了计数模块对密集小符号(如上下标、积分号)的感知能力。
2.3 Attention Decoder:带位置编码与覆盖率惩罚的序列解码
AttDecoder使用单层 GRU(GRUCell)逐符号自回归解码,其关键设计包括:
- PositionEmbeddingSine:对编码器特征叠加正弦位置编码,补偿 CNN 缺乏位置先验的问题;
- Coverage 注意力:
Attention模块将历史注意力累积alpha_sum通过卷积(kernel_size=11)与线性层映射,与当前隐藏状态、编码特征相加计算注意力分数,抑制注意力重复聚焦同一区域; - Counting 约束融合:计数向量经
counting_context_weight线性映射后,与隐藏状态、词嵌入、上下文向量求和,共同决定当前符号的输出分布word_prob,使解码过程"知道"每个符号应该出现几次。
训练时解码器按is_train=True使用教师强制(teacher forcing)逐位取标签labels[:, i]作为下一步输入;推理时(is_train=False)则取上一步argmax结果自回归生成,直至max_text_length(默认 36)结束。
2.4 损失函数与评估指标
- 损失函数:
CANLoss(见 rec_can_loss.py)由两部分组成——符号序列的CrossEntropyLoss(词级损失)与三个计数预测(counting_preds1、counting_preds2、取平均后的counting_preds)相对真实计数的SmoothL1Loss(计数损失)之和。真实计数标签由gen_counting_label按类别直方图生成,并忽略[0, 1, 107, 108, 109, 110]等特殊 token; - 评估指标:
CANMetric(见 rec_metric.py)基于SequenceMatcher计算字符级相似度,统计word_rate(符号级)与exp_rate(公式级)两个指标,配置文件中以main_indicator: exp_rate作为主指标。
3. 环境配置
在开始训练前,请先完成 PaddleOCR 运行环境的准备与项目代码的克隆:
- 运行环境准备参考《运行环境准备》;
- 项目代码克隆参考《项目克隆》。
CAN 模型的训练数据为 CROHME 数据集,官方以黑底白字(手写公式为白色、背景为黑色)的格式提供。训练数据目录结构需与配置文件保持一致,即./train_data/CROHME/training/(images + labels.txt)与./train_data/CROHME/evaluation/(images + labels.txt)。
4. 模型训练
PaddleOCR 对代码进行了模块化,训练 CAN 识别模型时需要更换配置文件为 rec_d28_can.yml。详细训练流程可参考文本识别训练教程。
4.1 启动训练
完成数据准备后即可启动训练:
# 单卡训练(训练周期长,不建议) python3 tools/train.py -c configs/rec/rec_d28_can.yml # 多卡训练,通过 --gpus 参数指定卡号 python3 -m paddle.distributed.launch --gpus '0,1,2,3' tools/train.py -c configs/rec/rec_d28_can.yml4.2 训练参数与注意事项
配置文件 rec_d28_can.yml 中几个关键训练参数:
| 配置项 | 默认值 | 说明 |
|---|---|---|
Global.epoch_num | 240 | 总训练轮数 |
Global.eval_batch_step | [0, 1105] | 每 1105 次 iteration(即 1 个 epoch,batch_size=8 时)评估一次 |
Global.character_dict_path | ppocr/utils/dict/latex_symbol_dict.txt | LaTeX 符号字典,CAN 专用 |
Global.max_text_length | 36 | 最大输出序列长度 |
Optimizer.name | Momentum | 动量优化器,momentum=0.9,clip_norm_global=100.0 |
Optimizer.lr | TwoStepCosine,lr=0.01,warmup_epoch=1 | 两段式余弦学习率衰减 |
Train.dataset.transforms | 含GrayImageChannelFormat: inverse: True | 黑底白字预处理(灰度图取反) |
Train.loader.batch_size_per_card | 8 | 单卡 batch size |
Train.loader.collate_fn | DyMaskCollator | 动态 mask 整理器,用于生成图像 mask 与标签 mask |
训练时需要特别注意以下两点:
图像颜色模式:官方提供的 CROHME 数据集将手写公式存储为黑底白字格式,因此配置中
GrayImageChannelFormat.inverse: True会在灰度化后取反图像。若您自行准备的数据集为白底黑字,请关闭取反:python3 tools/train.py -c configs/rec/rec_d28_can.yml -o Train.dataset.transforms.GrayImageChannelFormat.inverse=False评估频率:默认每训练 1 个 epoch(1105 次 iteration)评估 1 次,该值与
batch_size=8挂钩。若您更改 batch_size 或更换数据集,请按"数据集长度 // batch_size"重新计算并覆盖评估步数:python3 tools/train.py -c configs/rec/rec_d28_can.yml -o Global.eval_batch_step=[0, {length_of_dataset//batch_size}]
此外,标签编码由CANLabelEncode(见 label_ops.py)完成:它将 LaTeX 符号序列按空格分词,逐 token 映射为字典索引,并追加结束符</s>;字典中不存在的符号会被跳过。因此自备数据集时,务必保证标签序列中的符号全部存在于latext_symbol_dict.txt字典中。
5. 模型评估
可下载已训练完成的模型文件,使用如下命令进行评估:
# 注意将 pretrained_model 的路径设置为本地路径。 # 若使用自行训练保存的模型,请注意修改路径和文件名为 {path/to/weights}/{model_name}。 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评估过程复用配置文件中Eval段的数据集与预处理(同样含GrayImageChannelFormat.inverse: True),最终输出word_rate与exp_rate两项指标,其中exp_rate即文档表中所列的 51.72%(对应官方预训练模型在 CROHME 测试集上的表现)。
6. 模型预测
使用如下命令进行单张图片预测:
# 注意将 pretrained_model 的路径设置为本地路径。 python3 tools/infer_rec.py -c configs/rec/rec_d28_can.yml -o Architecture.Head.attdecoder.is_train=False Global.infer_img='./doc/datasets/crohme_demo/hme_00.jpg' Global.pretrained_model=./rec_d28_can_train/best_accuracy.pdparams # 预测文件夹下所有图像时,可修改 infer_img 为文件夹,如 Global.infer_img='./doc/datasets/crohme_demo/'。关键点说明:
Architecture.Head.attdecoder.is_train=False必须显式指定,使解码器切换为自回归推理模式(训练时为教师强制模式);- 预测的输入图像要求为黑底白字(与训练数据一致),手写公式为白色、背景为黑色;
- 若自行训练时修改过字典,需同步检查
Global.character_dict_path指向的字典文件是否正确。
7. 推理部署
7.1 导出 Inference 模型
首先将训练得到的最优模型转换成静态图 inference model。以官方训练完成的模型为例(模型下载地址):
# 注意将 pretrained_model 的路径设置为本地路径。 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_train=False # 目前的静态图模型默认的最大输出长度为 36, # 如果您需要预测更长的序列,请在导出模型时指定合适的输出长度,例如 Architecture.Head.max_text_length=72注意:如果您是在自己的数据集上训练的模型并调整了字典文件,请确认配置文件中的character_dict_path指向的是所需字典。
转换成功后,目录下会生成三个文件:
/inference/rec_d28_can/ ├── inference.pdiparams # 识别 inference 模型的参数文件 ├── inference.pdiparams.info # 识别 inference 模型的参数信息,可忽略 └── inference.pdmodel # 识别 inference 模型的 program 文件7.2 使用 predict_rec.py 推理
执行如下命令进行模型推理:
python3 tools/infer/predict_rec.py --image_dir="./doc/datasets/crohme_demo/hme_00.jpg" --rec_algorithm="CAN" --rec_batch_num=1 --rec_model_dir="./inference/rec_d28_can/" --rec_char_dict_path="./ppocr/utils/dict/latex_symbol_dict.txt" # 预测文件夹下所有图像时,可修改 image_dir 为文件夹,如 --image_dir='./doc/datasets/crohme_demo/'。 # 如果您需要在白底黑字的图片上进行预测,请设置 --rec_image_inverse=False在 predict_rec.py 中,当rec_algorithm == "CAN"时:
- 后处理选用
CANLabelDecode(见 rec_postprocess.py),它沿时间维取argmax得到符号索引序列,以序列中第一个结束符位置截断,再将索引逐项映射回 LaTeX 符号并以空格连接输出; - 预处理调用
norm_img_can:先将图像转为灰度图,若rec_image_inverse=True(默认)则执行255 - img取反,再按(1, 32, 320)形状进行等比缩放与填充(见 predict_rec.py 中 norm_img_can 实现); - 推理输入为
[norm_img_batch, norm_img_mask_batch, word_label_list]三元组,其中 mask 全 1、标签为全 1 的占位序列,与训练阶段的多输入结构保持一致。
对上方示例图片执行命令后,预测结果(识别的 LaTeX 符号序列)会打印到屏幕上:
Predicts of ./doc/imgs_hme/hme_00.jpg:['x _ { k } x x _ { k } + y _ { k } y x _ { k }', []]推理注意事项:
- 预测图像必须为黑底白字(手写公式为白色、背景为黑色);
- 推理时需通过
rec_char_dict_path指定字典,若您修改了字典,请同步修改该参数; - 若您修改了预处理方法,需修改 predict_rec.py 中 CAN 的预处理为您的预处理方法。
7.3 C++ / Serving / 更多推理部署
由于 C++ 预处理与后处理尚未支持 CAN,C++ 推理部署暂未支持;Serving 服务化部署与更多推理部署(如 Paddle Lite 等)当前同样暂不支持。该限制明确记载于 algorithm_rec_can.md 文档中,部署到上述平台前请留意此约束。
8. FAQ
- CROHME 数据集从何而来?CROHME 数据集来自于 CAN 源 repo(https://github.com/LBH1024/CAN),PaddleOCR 在 rec_can_head.py 与 rec_can_loss.py 的代码注释中也明确标注了参考来源。
- 为什么 CAN 需要专门的字典?CAN 的输出是 LaTeX 符号 token 序列而非普通文本,因此必须使用 latext_symbol_dict.txt(111 类,含
</s>等特殊 token),不能复用通用中英文字典。
9. 引用
@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} }【免费下载链接】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),仅供参考