尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

CAN模型实战:手写数学公式识别从理论到部署全解析

CAN模型实战:手写数学公式识别从理论到部署全解析 1. 项目概述从论文到实践复现CAN模型最近在整理手写数学公式识别的开源项目发现很多朋友对2021年发表在CVPR上的那篇《CAN: Counting-Aware Network for Handwritten Mathematical Expression Recognition》很感兴趣。这篇论文提出的CAN模型通过引入一个计数感知模块来辅助识别在当时多个公开数据集上刷出了SOTAState-of-the-the-Art结果。网上能找到的官方代码仓库对于刚接触这个领域的研究者或开发者来说阅读和运行起来可能有些门槛。我自己花了些时间把代码从头到尾梳理了一遍并且成功用自己的数据集完成了训练和评估。这篇文章我就来分享一下整个过程的详细步骤、踩过的坑以及一些实用的调参经验目标是让你能拿着这份“攻略”快速复现论文结果并应用到自己的数据上。手写数学公式识别HMER这个任务本质上是一个结合了图像识别和序列生成的视觉-语言任务。它比一般的OCR要复杂得多因为公式具有二维结构比如上下标、分式、根号并且符号间存在复杂的语法关系。CAN模型的创新点在于它不仅仅依赖编码器-解码器框架去“猜”下一个符号还额外引入了一个“计数”分支。这个分支会预测公式图像中每个数学符号出现的次数作为一个全局的上下文信息来约束和指导解码过程从而减少符号重复或遗漏的错误。这个思路非常直观有效尤其是在处理长公式或者包含大量相同符号比如一连串的“a”或“1”的公式时。如果你正在做相关研究或者需要在自己的业务场景比如教育领域的作业批改、科研笔记数字化中集成公式识别功能那么理解并实践CAN模型会是一个很好的起点。接下来我会从环境搭建、代码结构解析、数据准备、训练调试到最终推理一步步拆开来讲。2. 核心思路与代码结构深度解析2.1 CAN模型的核心思想为什么“计数”能提升识别率在深入代码之前我们必须先吃透论文的核心思想。传统的基于注意力机制的编码器-解码器模型比如Show-Attend-Tell及其在HMER上的变种在解码时解码器严重依赖于编码器提供的视觉特征和上一步生成的符号。这种方式有时会陷入局部最优比如重复生成某个符号或者漏掉某个该出现的符号。CAN的作者认为一个公式中每个符号出现的次数是一个有价值的全局信息。例如知道图像里大概有3个“x”和2个“”那么解码器在生成时就会受到这个全局数量的“软约束”。具体实现上CAN在主干网络如DenseNet提取视觉特征后并行地接了两个头视觉计数模块这个模块不是直接数数而是通过一个回归网络从视觉特征中预测出一个“计数向量”。这个向量的长度等于词汇表大小每个位置的值代表对应符号的预测出现次数是一个连续值不是整数。识别主干这就是传统的编码器-解码器编码器通常是一个CNN如DenseNet加位置编码解码器是一个基于注意力机制的LSTM或Transformer。在训练时模型有两个损失函数计数损失预测的计数向量与真实符号计数的均方误差MSE。识别损失解码器生成的符号序列与真实序列的交叉熵损失Cross-Entropy。最终的总损失是两者的加权和。在推理预测时计数模块的预测结果会被转换成一种先验知识融入到解码器的初始化状态或每一步的上下文计算中从而引导解码过程。这种“计数感知”的机制相当于给模型增加了一个全局的校验器有效提升了识别的准确率特别是对于复杂的长公式。2.2 官方代码仓库结构梳理官方代码通常托管在GitHub上。我们以典型的PyTorch实现为例来梳理其目录结构。理解这个结构是后续一切操作的基础。CAN-HMER/ ├── config/ # 配置文件目录 │ └── can.yml # 模型超参数、路径等配置 ├── dataset/ # 数据集相关 │ ├── __init__.py │ ├── hme_dataset.py # 核心数据集加载类 │ └── utils.py # 数据预处理工具 ├── models/ # 模型定义 │ ├── __init__.py │ ├── can.py # CAN模型主类 │ ├── decoder.py # 解码器如AttnDecoder │ ├── encoder.py # 编码器如DenseNetEncoder │ └── counting.py # 计数模块 ├── utils/ # 工具函数 │ ├── metrics.py # 评估指标ExpRate, BLEU等 │ ├── tokenizer.py # 标签分词器将LaTeX序列转为id │ └── visualization.py # 注意力权重可视化 ├── train.py # 训练脚本主入口 ├── test.py # 测试/评估脚本 ├── predict.py # 单张图片预测脚本 ├── requirements.txt # Python依赖包列表 └── README.md # 项目说明关键文件解读config/can.yml这是项目的控制中枢。你需要在这里修改数据路径、模型结构编码器类型、解码器维度、训练参数学习率、batch_size、计数损失的权重等。第一次跑通后大部分调参工作都是通过修改这个文件完成的。dataset/hme_dataset.py定义了如何读取图片和对应的LaTeX标签文件以及进行了哪些数据增强如随机裁剪、缩放、归一化。这是适配自己数据集时需要重点修改的文件。models/can.py这里是CAN模型的整体架构它整合了编码器、计数模块和解码器并定义了前向传播和损失计算的过程。train.py训练循环的逻辑。包括加载数据、模型、优化器以及每个epoch的训练和验证步骤保存最佳模型等。注意不同研究者复现的CAN代码可能在细节上有差异例如解码器可能用LSTMAttention也可能用Transformer。但核心的多任务损失框架和计数模块的集成方式是统一的。在开始之前务必通读一遍train.py和models/can.py理解数据流和损失计算的具体位置。3. 环境搭建与数据准备实战3.1 构建可复现的Python环境深度学习项目的第一道坎往往是环境。为了避免“在我机器上能跑”的尴尬强烈建议使用Conda进行环境管理。# 1. 创建并激活一个全新的conda环境以Python 3.8为例兼容性较好 conda create -n can-hmer python3.8 -y conda activate can-hmer # 2. 根据项目提供的requirements.txt安装核心依赖 # 通常包含pytorch, torchvision, opencv-python, nltk, pyyaml, tensorboard等 pip install -r requirements.txt # 3. 重点PyTorch的安装需要去官网根据你的CUDA版本选择命令 # 例如对于CUDA 11.3 pip install torch1.12.1cu113 torchvision0.13.1cu113 --extra-index-url https://download.pytorch.org/whl/cu113 # 4. 安装LaTeX渲染相关包用于可视化或后处理可选但推荐 pip install Pillow matplotlib实操心得如果项目没有提供requirements.txt你可以通过pip freeze requirements.txt在原作者的环境中生成或者根据代码中的import语句手动安装。常见的必有包torch,torchvision,opencv-python,nltk用于BLEU评分,tensorboard或wandb用于日志记录。PyTorch版本与CUDA版本的匹配至关重要。使用nvidia-smi查看驱动支持的CUDA最高版本然后去 PyTorch官网 查找对应的安装命令。版本不匹配会导致无法调用GPU。遇到“No module named ‘utils.*’”这类错误通常是因为Python的模块导入路径问题。确保在项目根目录下运行脚本或者将项目路径添加到PYTHONPATH环境变量中。3.2 准备自己的数据集格式与预处理CAN论文主要在CROHME手写公式识别竞赛数据集上评估。但我们要用自己的数据训练就必须将数据整理成模型能接受的格式。1. 数据格式要求模型通常需要两种文件图像文件公式的灰度图或二值图。建议统一处理为灰度图尺寸不定但长宽最好能归一化到某个范围如高度固定为64宽度按比例缩放。标签文件一个文本文件如train_label.txt每一行对应一张图片的标签。格式为图片路径\t LaTeX表达式。data/train/img_001.png x_{1} y^{2} \\frac{a}{b} data/train/img_002.png \\sum_{i1}^{n} i \\frac{n(n1)}{2}注意LaTeX表达式中的反斜杠\需要转义写成\\。这是为了在Python读取文本文件时能正确保留。2. 创建词汇表模型需要一个词汇表文件vocab.txt包含所有可能出现的符号token。这通常包括数字0-9字母a-z, A-Z希腊字母\alpha, \beta, ...运算符, -, \times, \div, , , , ...结构符号_{, }, ^{, }, \frac{, }{, }, \sqrt{, }, ...特殊标记sos序列开始,eos序列结束,pad填充,unk未知符号你可以编写一个脚本遍历所有标签文件用正则表达式匹配LaTeX命令和特殊字符或简单的空格分割如果标签已预处理为token序列来提取所有独特的符号然后生成这个文件。3. 修改数据集加载代码这是最关键的一步。打开dataset/hme_dataset.py找到__getitem__方法。你需要确保它能根据标签文件中的路径正确读取你的图片。对你的图片进行了与原始CROHME数据类似的预处理如转为灰度、归一化、尺寸调整、数据增强。调用tokenizer将LaTeX标签字符串转换为索引id序列并自动添加sos和eos标记。常见问题与处理技巧图片尺寸差异大直接resize到固定高、变宽可能会严重变形。常用做法是固定高度如64宽度按原图比例缩放然后将宽度填充pad到某个固定值如256或当前batch内的最大宽度。在hme_dataset.py的collate_fn函数中处理填充。LaTeX语法不规范自己标注的数据可能存在语法错误或不统一如有时用\frac有时用\dfrac。建议在生成词汇表前对标签进行清洗和标准化。可以定义一个LaTeX命令的映射表将不同写法统一。数据量不足手写公式数据标注成本高。如果数据量少10k模型很容易过拟合。除了使用Dropout、权重衰减等正则化方法可以尝试强数据增强对图片进行弹性变换、随机涂抹、添加高斯噪声等模拟不同的书写风格和噪声。预训练在公开的大规模数据集如CROHME、HME100K上先预训练模型再用自己的小数据集进行微调Fine-tuning。这是提升小数据集性能最有效的手段之一。4. 模型训练全流程与调参详解4.1 配置文件修改与训练启动一切准备就绪后我们开始训练。首先根据你的数据和硬件调整config/can.yml。# config/can.yml 关键部分示例 data: train_label: ‘path/to/your/train_label.txt‘ # 修改为你的路径 eval_label: ‘path/to/your/val_label.txt‘ # 修改为你的路径 vocab: ‘path/to/your/vocab.txt‘ # 修改为你的路径 image_height: 64 # 图片固定高度 image_width: 256 # 图片最大宽度/填充宽度 model: encoder: ‘densenet121‘ # 编码器类型 decoder: ‘attn_lstm‘ # 解码器类型 decoder_dim: 512 # 解码器隐藏层维度 attention_dim: 512 # 注意力维度 counting_dim: 512 # 计数模块维度 dropout: 0.3 # Dropout率防过拟合 training: batch_size: 16 # 根据GPU内存调整越大越稳但耗内存 epochs: 100 # 训练轮数 learning_rate: 1.0 # 初始学习率对于Adam优化器可能偏高 optimizer: ‘adam‘ # 优化器 lr_scheduler: ‘step‘ # 学习率调度器 step_size: 30 # 每30个epoch学习率衰减 gamma: 0.1 # 衰减系数 counting_loss_weight: 1.0 # 计数损失的权重重要超参 logging: log_dir: ‘runs/exp1‘ # TensorBoard日志目录 save_dir: ‘saved_models/exp1‘ # 模型保存目录修改完毕后在终端运行训练命令python train.py --config config/can.yml4.2 训练过程监控与问题诊断训练开始后不要干等着。要通过日志和可视化工具密切监控。使用TensorBoardtensorboard --logdir runs/exp1 --port 6006然后在浏览器打开localhost:6006。重点关注以下曲线损失曲线train_loss和val_loss。理想情况是两者同步下降且val_loss在后期平稳或缓慢上升可能过拟合。如果train_loss下降但val_loss很早就开始飙升是典型的过拟合。识别准确率train_exp_rate和val_exp_rateExpRate即完全匹配准确率。这是我们的核心指标。计数损失train_count_loss和val_count_loss。观察它是否在正常下降如果一直很高可能是计数模块设计或损失权重有问题。常见训练问题与调参策略问题损失Loss不下降或为NaN。检查学习率这是最常见的原因。对于Adam优化器论文中可能用1.0但这对于很多任务来说太高了。我个人的经验是从一个较小的值开始尝试比如3e-4或1e-3。如果损失爆炸变成NaN立即停止训练降低学习率10倍再试。检查梯度可以在train.py中添加梯度裁剪torch.nn.utils.clip_grad_norm_防止梯度爆炸。检查数据确认数据加载是否正确有没有损坏的图片或无法解析的标签。可以在数据集类的__getitem__方法中加入简单的打印或断言来调试。检查损失函数确认计数损失MSE的数值范围是否合理。如果计数标签是很大的整数MSE可能会非常大导致总损失被主导。可以考虑对计数标签进行归一化或者调整counting_loss_weight先尝试设为0.1或0.01。问题训练集准确率很高但验证集准确率很低过拟合。增加正则化增大dropout率如从0.3调到0.5在编码器和解码器中都应用Dropout。使用权重衰减在优化器中加入weight_decay参数如1e-5。加强数据增强在hme_dataset.py的预处理部分添加更多样化的增强如随机旋转小角度、透视变换、对比度调整等。减少模型容量如果数据量很小可以尝试使用更小的编码器如DenseNet-121换成更小的网络或者减少decoder_dim。早停Early Stopping监控val_exp_rate当其在连续多个epoch如10个不再提升时停止训练并回滚到最佳模型。问题计数损失下降很慢或者对最终识别准确率提升不明显。调整计数损失权重counting_loss_weight是一个关键的超参数。如果权重太大模型会过于关注计数任务而忽略主识别任务如果太小则计数模块起不到作用。建议的做法是进行网格搜索比如尝试[0.01, 0.1, 0.5, 1.0, 2.0]观察哪个值在验证集上能获得最高的识别准确率。检查计数标签确认你为每张图片生成的“真实计数向量”是否正确。计数模块学习的是一个回归任务如果标签有误它就无法学到有用的信息。审视计数模块结构论文中的计数模块可能是一个简单的多层感知机MLP。如果问题复杂可以尝试加深或加宽这个网络或者引入更复杂的结构如基于注意力的计数。4.3 模型评估与指标解读训练完成后使用test.py脚本在独立的测试集上评估模型性能。python test.py --config config/can.yml --checkpoint saved_models/exp1/best_model.pth --eval_split test关键评估指标表达式识别率ExpRate这是最严格的指标要求预测的整个LaTeX序列与真实序列完全一致包括所有空格和符号。这是论文报告的主要指标。BLEU分数来自机器翻译的指标衡量预测序列和真实序列在n-gram上的重合度。它比ExpRate宽松能部分反映语义相似性。但要注意对于公式这种结构严谨的序列BLEU高不一定代表公式正确。树编辑距离TED一种更符合公式结构的指标它先将LaTeX序列解析成语法树然后计算两棵树之间的编辑距离。这个指标更能反映结构错误。但实现起来较复杂不是所有开源代码都包含。如何解读结果如果你的模型在自己测试集上的ExpRate达到80%以上说明模型已经学习得相当不错了。与原始论文在CROHME上的结果如CROHME 2014上ExpRate约56%对比时要谨慎。因为数据集不同书写风格、复杂度、词汇表直接比较数字意义不大。更重要的是看模型在你关心的业务场景下的实际效果。分析错误案例随机抽样一些识别错误的样本观察是哪些类型的错误符号混淆、结构错误、多符、漏符。这能为你下一步的改进如调整数据增强、修改模型提供最直接的线索。5. 推理部署与性能优化思考5.1 单张图片预测与可视化模型训练好后我们可以用predict.py脚本或自己写一个简单的推理脚本来识别单张图片。import torch from PIL import Image from models.can import CAN from utils.tokenizer import Tokenizer import yaml from dataset import build_preprocess_transform # 1. 加载配置和模型 with open(‘config/can.yml‘, ‘r‘) as f: config yaml.safe_load(f) checkpoint torch.load(‘saved_models/exp1/best_model.pth‘, map_location‘cpu‘) model CAN(config[‘model‘]) model.load_state_dict(checkpoint[‘model_state_dict‘]) model.eval() # 2. 加载词汇表和分词器 tokenizer Tokenizer(config[‘data‘][‘vocab‘]) # 3. 预处理图片 transform build_preprocess_transform(config[‘data‘][‘image_height‘], config[‘data‘][‘image_width‘]) image Image.open(‘your_formula.png‘).convert(‘L‘) # 转为灰度 image_tensor transform(image).unsqueeze(0) # 增加batch维度 # 4. 预测 with torch.no_grad(): pred, _ model(image_tensor, mode‘eval‘) # 使用eval模式不计算损失 pred_seq pred[0].argmax(dim-1).cpu().numpy() # 获取概率最大的token id序列 # 5. 解码为LaTeX字符串 latex_str tokenizer.decode(pred_seq, remove_special_tokensTrue) print(‘Predicted LaTeX:‘, latex_str)可视化注意力权重一个很有用的调试工具是可视化解码过程中的注意力权重。这能帮你理解模型在生成每个符号时“看”了图片的哪个区域。如果发现注意力散乱或与符号位置不对齐可能意味着编码器特征提取有问题或者注意力机制需要调整。相关代码通常在utils/visualization.py中。5.2 模型优化与加速部署考量如果要将模型投入实际应用还需要考虑性能和效率。模型轻量化更换轻量编码器DenseNet虽然性能好但参数量和计算量较大。可以考虑替换为MobileNetV3、EfficientNet-Lite或GhostNet等轻量级网络并在你的数据上重新微调。知识蒸馏用一个大的、训练好的CAN模型教师模型去指导一个小的学生模型训练在尽量不损失精度的情况下减少模型尺寸。剪枝与量化移除模型中不重要的连接剪枝并将权重从FP32转换为INT8量化。PyTorch提供了相关的工具如torch.quantization但这通常会带来一定的精度损失需要仔细评估。推理加速TorchScript将PyTorch模型转换为TorchScript可以获得更快的推理速度并且易于在C环境中部署。ONNX Runtime / TensorRT将模型导出为ONNX格式然后利用ONNX Runtime或NVIDIA TensorRT进行高性能推理尤其能充分发挥GPU的潜力。批处理Batch Inference在实际服务中对输入的多个图片进行批处理能极大提升GPU的利用率和吞吐量。确保你的推理脚本支持批处理。错误后处理 模型预测的LaTeX序列可能包含一些语法错误。可以设计一个简单的后处理规则例如检查括号是否匹配。将连续的、相同的基础符号如x x x合并为x_{3}如果计数模块支持的话这个信息可以来自计数模块的预测。使用一个简单的LaTeX语法检查器如果存在来纠正明显错误。从读论文、捋代码到配环境、改数据、跑训练、调参数最后看到模型能正确识别出自己手写的公式这个过程虽然繁琐但成就感十足。CAN模型将计数信息引入识别框架的思路非常巧妙它启发了我们在解决序列生成问题时除了局部注意力全局的、结构化的先验知识能起到强大的约束作用。在实际操作中最大的挑战往往不是模型本身而是数据的准备和清洗以及训练过程中那些“玄学”般的超参数调试。我的经验是保持耐心从一个小而稳定的配置开始比如较低的学习率、较强的数据增强建立baseline然后每次只调整一个变量并详细记录实验日志这样才能逐步逼近最优解。希望这份详细的梳理和实战记录能帮你少走些弯路。
返回列表