1. 项目概述:从论文到实践,复现CAN模型
最近在整理手写数学公式识别的开源项目,发现很多朋友对2021年发表在CVPR上的那篇《CAN: Counting-Aware Network for Handwritten Mathematical Expression Recognition》很感兴趣。这篇论文提出的CAN模型,通过引入一个计数感知模块来辅助识别,在当时多个公开数据集上刷出了SOTA(State-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代码可能在细节上有差异,例如解码器可能用LSTM+Attention,也可能用Transformer。但核心的多任务损失框架和计数模块的集成方式是统一的。在开始之前,务必通读一遍
train.py和models/can.py,理解数据流和损失计算的具体位置。
3. 环境搭建与数据准备实战
3.1 构建可复现的Python环境
深度学习项目的第一道坎往往是环境。为了避免“在我机器上能跑”的尴尬,强烈建议使用Conda进行环境管理。
# 1. 创建并激活一个全新的conda环境(以Python 3.8为例,兼容性较好) conda create -n can-hmer python=3.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 torch==1.12.1+cu113 torchvision==0.13.1+cu113 --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表达式。
注意:LaTeX表达式中的反斜杠data/train/img_001.png x_{1} + y^{2} = \\frac{a}{b} data/train/img_002.png \\sum_{i=1}^{n} i = \\frac{n(n+1)}{2}\需要转义,写成\\。这是为了在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 训练过程监控与问题诊断
训练开始后,不要干等着。要通过日志和可视化工具密切监控。
使用TensorBoard:
tensorboard --logdir runs/exp1 --port 6006然后在浏览器打开
localhost:6006。重点关注以下曲线:- 损失曲线:
train_loss和val_loss。理想情况是两者同步下降,且val_loss在后期平稳或缓慢上升(可能过拟合)。如果train_loss下降但val_loss很早就开始飙升,是典型的过拟合。 - 识别准确率:
train_exp_rate和val_exp_rate(ExpRate,即完全匹配准确率)。这是我们的核心指标。 - 计数损失:
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_tokens=True) 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,然后每次只调整一个变量,并详细记录实验日志,这样才能逐步逼近最优解。希望这份详细的梳理和实战记录,能帮你少走些弯路。