ARTICLE DETAIL

资讯详情

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

从零构建手写公式识别引擎:CAN模型实战与数据集训练指南

从零构建手写公式识别引擎:CAN模型实战与数据集训练指南 1. 项目概述从零到一构建自己的手写公式识别引擎最近在整理手写数学公式识别的相关技术发现基于CANComposition Attention Network模型的论文依然是这个领域一个非常扎实且经典的工作。很多朋友拿到开源代码后面对复杂的项目结构和论文中的数学符号往往不知从何下手。更有甚者想要用自己的数据集比如特定学科的手写笔记、工程计算草稿来训练模型却卡在了数据预处理和训练流程上。这篇文章我就结合自己复现和改造CAN模型的经验带大家彻底梳理一遍代码并手把手教你如何准备和训练自己的数据集。整个过程我会尽量避开那些“正确的废话”直接分享实操中会遇到的关键步骤和踩过的坑。手写数学公式识别Handwritten Mathematical Expression Recognition, HMER本身是个很有挑战性的任务它结合了手写文字识别和二维结构分析。CAN模型的核心创新在于其“组合注意力”机制它不像传统方法那样简单地按行或按列识别而是试图理解公式中符号之间的空间布局关系比如上下标、分式、根号等。这对于准确还原LaTeX序列至关重要。我们接下来的目标很明确第一理解CAN代码的每一部分在干什么第二学会准备符合模型要求的数据格式第三成功跑通训练并评估自己的模型。2. CAN模型核心思想与代码架构拆解在深入代码之前我们必须先搞懂CAN模型到底在解决什么问题以及它是如何解决的。这能帮助我们在看代码时不是盲人摸象而是心中有图。2.1 问题本质与模型输入输出手写公式识别本质上是一个“图像到序列”的翻译问题。输入是一张包含手写数学公式的灰度图片输出是一个序列通常是LaTeX字符串用于描述这个公式。例如图片上是一个手写的分数“½”模型应该输出“\frac{1}{2}”。这个任务的难点在于二维空间结构公式符号不是线性排列的。下标、上标、分式的分子分母等构成了复杂的二维关系。符号多样性数学符号种类繁多包括字母、数字、运算符、希腊字母、特殊符号等。手写变体不同人的笔迹差异大同一人的书写也有波动。CAN模型通过编码器-解码器Encoder-Decoder框架来解决这个问题。编码器通常是一个CNN如DenseNet负责从图像中提取视觉特征。解码器一个基于注意力机制的RNN如LSTM负责根据这些特征一步步生成输出序列。2.2 “组合注意力”机制的精髓CAN的核心——“组合注意力”Composition Attention是它区别于普通注意力机制的关键。普通注意力如Bahdanau Attention在解码的每一步会计算解码器当前状态与编码器所有特征之间的相关性得到一个权重向量然后用这个权重对编码器特征进行加权求和得到一个“上下文向量”。CAN的组合注意力在此基础上做了升级。它认为在生成公式的某个符号时比如生成分式的横线“-”模型不仅需要关注这个符号对应的图像区域还需要同时关注与之有结构关系的其他区域比如分子和分母的位置。因此CAN的注意力模块会输出多个注意力权重图每个图聚焦于与当前生成符号相关的不同视觉组件然后将这些组件的信息组合起来供解码器使用。在代码中这通常体现为一个名为CompositionAttention的模块。它会接收解码器的隐藏状态和编码器的特征图然后并行计算K个论文中K4注意力权重图alpha_k最后将这些权重图分别与特征图加权求和得到K个组件向量再通过一个小的神经网络组合成最终的上下文向量。# 伪代码示意帮助理解流程 class CompositionAttention(nn.Module): def __init__(self, decoder_dim, encoder_dim, attention_dim, K4): super().__init__() self.K K # 用于计算K个注意力权重的线性层 self.attention_layers nn.ModuleList([... for _ in range(K)]) self.combine nn.Linear(K * encoder_dim, encoder_dim) # 组合组件向量的网络 def forward(self, decoder_hidden, encoder_features): components [] for k in range(self.K): # 计算第k个注意力权重图 alpha_k energy_k self.attention_layers[k](decoder_hidden, encoder_features) alpha_k F.softmax(energy_k, dim1) # 加权求和得到第k个组件向量 context_k context_k (encoder_features * alpha_k.unsqueeze(2)).sum(dim1) components.append(context_k) # 拼接所有组件向量并组合 combined_context torch.cat(components, dim1) final_context self.combine(combined_context) return final_context, alpha_stack # 返回最终的上下文向量和注意力图用于可视化2.3 代码仓库结构梳理一个典型的CAN开源实现例如基于PyTorch的目录结构可能如下所示。理解这个结构是上手的第一步CAN-HMER/ ├── configs/ # 配置文件目录 │ └── can.yml # 模型超参数、路径等配置 ├── data/ # 数据相关 │ ├── CROHME/ # 标准数据集如CROHME存放处 │ │ ├── train_images/ │ │ ├── train_labels.txt │ │ └── ... │ └── preprocess.py # 数据预处理脚本 ├── datasets/ # PyTorch Dataset类定义 │ └── crohme_dataset.py # 加载和解析数据 ├── models/ # 模型定义 │ ├── encoder.py # 编码器CNN │ ├── decoder.py # 解码器带组合注意力的RNN │ └── can.py # 整合编码器-解码器的CAN主模型 ├── utils/ # 工具函数 │ ├── metrics.py # 评估指标如ExpRate, WER │ ├── vocabulary.py # 词表构建与管理 │ └── visualization.py # 可视化注意力图 ├── train.py # 模型训练主脚本 ├── eval.py # 模型评估脚本 ├── predict.py # 单张图片预测脚本 └── requirements.txt # 项目依赖关键文件解读configs/can.yml这是项目的控制中心。所有重要参数都在这里设置如图像尺寸、批次大小、学习率、编码器/解码器维度、词表路径、数据路径等。修改配置是适配自己数据集的首要步骤。datasets/crohme_dataset.py定义了如何读取一张图片和其对应的LaTeX标签。你需要重点修改这里的__getitem__方法使其兼容你自定义的数据格式。models/can.py这里是CAN模型的完整定义它实例化了编码器和解码器并定义了前向传播的逻辑。通常不需要大改除非你想调整模型结构。train.py训练循环。包含了损失计算通常是交叉熵损失、优化器如Adam、学习率调度、模型保存等逻辑。你需要关注数据加载器是如何构建的。注意不同的开源实现结构可能有差异但核心模块Encoder, Decoder, Attention, Dataset, Trainer是共通的。我们的目标是找到这些核心部分并理解它们之间的数据流。3. 准备自己的数据集从原始图片到模型可读格式用自己的数据训练是本文的重点也是难点。公开数据集如CROHME已经提供了规整的图片和标签但我们自己的数据往往是杂乱无章的。这个过程比跑通官方代码更需要耐心和细心。3.1 数据收集与初步整理假设你有一批手写数学公式的图片可能来自扫描的笔记、平板电脑的手写记录或者拍照的草稿纸。第一步是整理统一格式将所有图片转换为同一种格式推荐.png或.jpg。确保色彩模式为灰度图单通道这通常是模型输入的要求。你可以用PIL或OpenCV批量转换。统一命名建议使用有规律的命名如formula_001.png,formula_002.png... 这便于后续制作标签文件。基本清洗剔除模糊不清、过于倾斜或包含大量非公式内容的图片。可以简单目检如果数据量大可以考虑写脚本用图像清晰度指标进行初筛。3.2 标注工具与LaTeX标签生成这是最耗时的一步。你需要为每张图片生成对应的LaTeX序列标签。手动标注对于小规模数据集几百张你可以使用文本编辑器手动编写LaTeX。但这要求你对LaTeX数学语法非常熟悉且容易出错。半自动标注利用现有识别工具可以先用一些开源的或在线的公式识别工具如Mathpix Snip、LaTeX-OCR等对图片进行初步识别生成一个大概的LaTeX然后人工进行校对和修正。这能极大提升效率。构建标注界面如果需要持续标注可以考虑用Python的Tkinter或Web技术如FlaskHTML写一个简单的标注工具显示图片并提供文本框用于输入/修改LaTeX。无论用哪种方式最终你需要得到一个标签文件通常是一个.txt或.json文件。每行对应一张图片包含图片文件名和LaTeX标签中间用制表符或空格分隔。标签文件示例 (train_labels.txt):formula_001.png \frac { 1 } { 2 } \sqrt { x } formula_002.png y \sum _ { i 1 } ^ { n } a_i x^i formula_003.png \int _ { 0 } ^ { \infty } e^{-x^2} dx \frac{\sqrt{\pi}}{2}重要提示LaTeX标签中的空格处理很关键。许多实现要求像上面例子一样在符号和花括号周围加入空格进行分词Tokenization例如\frac { 1 } { 2 }而不是\frac{1}{2}。这取决于代码中vocabulary.py的分词方式你必须与训练代码保持一致。通常开源代码会提供一个tokenize函数你需要用同样的方式处理你的标签。3.3 数据预处理与Dataset类适配CAN模型对输入图片有固定要求比如尺寸、归一化等。你需要修改或确保预处理流程适用于你的数据。图像预处理常见的预处理包括调整大小将图片缩放到固定高度如100像素宽度按比例缩放或者直接缩放到固定尺寸如[H, W] [100, 300]。归一化将像素值从[0, 255]归一化到[0, 1]或[-1, 1]。数据增强为了提升模型泛化能力可以在训练时加入随机增强如小幅度的旋转、缩放、平移、弹性形变、添加噪声等。注意对于数学公式过度增强如大角度旋转可能会破坏其空间结构需谨慎使用。修改Dataset类这是连接你的数据和模型的桥梁。你需要打开datasets/crohme_dataset.py或类似文件重点关注__init__,__getitem__,__len__这三个方法。在__init__中读取你制作的标签文件将图片路径和标签字符串加载到两个列表中。在__getitem__中根据索引idx读取图片应用上述预处理变换同时将标签字符串转换为索引序列这个过程通常由vocabulary类完成。# 自定义Dataset类核心部分示例 from torch.utils.data import Dataset from PIL import Image import torchvision.transforms as transforms class MyFormulaDataset(Dataset): def __init__(self, label_file, img_dir, vocab, transformNone): self.img_dir img_dir self.vocab vocab self.transform transform self.samples [] with open(label_file, r, encodingutf-8) as f: for line in f: parts line.strip().split(\t) # 假设用制表符分隔 if len(parts) 2: img_name, latex parts self.samples.append((img_name, latex)) def __getitem__(self, idx): img_name, latex self.samples[idx] img_path os.path.join(self.img_dir, img_name) # 读取图片确保是灰度图 image Image.open(img_path).convert(L) if self.transform: image self.transform(image) # 将LaTeX字符串转换为单词索引序列 # 假设vocab有一个tokenize方法和一个words2indices方法 tokens self.vocab.tokenize(latex) target self.vocab.words2indices(tokens) target torch.LongTensor(target) return image, target def __len__(self): return len(self.samples)构建词表词表Vocabulary是模型认识的所有“单词”的集合。在公式识别中“单词”就是LaTeX符号如\frac,{,},x,1,等。你需要基于你的训练集标签来构建词表。通常代码中的utils/vocabulary.py会有一个build_vocab函数它遍历所有标签统计符号频率并建立符号到索引的映射。记得将构建好的词表通常是一个pickle文件保存下来并在配置文件中指定路径。4. 训练流程详解与关键参数调优当数据和代码都准备好后就可以开始训练了。这里我们深入train.py脚本看看每一步在做什么以及有哪些可以调整的“旋钮”。4.1 训练脚本核心循环解析一个标准的训练循环包含以下步骤初始化加载配置、创建模型、定义损失函数nn.CrossEntropyLoss、优化器torch.optim.Adam、学习率调度器。数据加载使用自定义的Dataset和DataLoader加载训练集和验证集。DataLoader的collate_fn函数需要特别注意因为公式图片的宽度不同标签序列长度也不同需要进行填充Padding以使一个批次内的数据形状一致。训练循环for epoch in range(num_epochs): model.train() for batch_idx, (images, targets) in enumerate(train_loader): optimizer.zero_grad() # 前向传播 outputs model(images, targets[:, :-1]) # 解码器输入需要去掉最后一个token # 计算损失忽略填充位置padding_idx loss criterion(outputs.view(-1, vocab_size), targets[:, 1:].contiguous().view(-1)) # 反向传播 loss.backward() # 梯度裁剪防止梯度爆炸RNN/Transformer常见操作 torch.nn.utils.clip_grad_norm_(model.parameters(), clip_grad) optimizer.step() # ... 记录日志 ... # 每个epoch后在验证集上评估 val_loss, val_accuracy evaluate(model, val_loader, criterion) # 根据验证集表现保存最佳模型调整学习率 scheduler.step(val_loss)4.2 关键超参数经验谈配置文件里的参数不是随便设的每个都影响着训练速度和最终效果。图像尺寸 (img_h,img_w)这是输入图片的高度和宽度。不是越大越好。更大的尺寸意味着编码器CNN需要处理更多的像素计算量剧增但识别精度可能不会线性提升。对于大多数手写公式[100, 300]或[128, 384]是一个不错的起点。可以先尝试用较小的尺寸快速实验。编码器维度 (encoder_dim)这是CNN提取出的特征图的通道数。它决定了视觉特征的丰富程度。DenseNet-121通常输出1024维的特征但后面可能会接一个1x1卷积来降维到encoder_dim如512。更大的维度能容纳更多信息但也更耗内存。解码器维度 (decoder_dim)这是LSTM隐藏状态的大小。它决定了解码器“记忆”和“思考”的能力。通常设置为256或512。需要与encoder_dim匹配因为注意力机制会连接两者。注意力组件数 (K)CAN论文中的关键参数表示组合注意力的组件数量。论文默认是4。理论上K越大模型捕捉不同结构组件的能力越强但参数和计算量也增加。不建议一开始就修改它先用默认值4。学习率 (learning_rate)这是最重要的参数之一。对于Adam优化器常见的初始学习率是1e-4或3e-4。如果训练初期损失不下降可以尝试调大如1e-3如果损失震荡剧烈可以调小。配合学习率调度器如ReduceLROnPlateau使用效果更好。批次大小 (batch_size)受限于GPU内存。在内存允许的情况下较大的批次大小如16, 32能使梯度估计更稳定可能有助于收敛。但如果数据差异大小批次如8有时能带来更好的泛化性能。需要根据你的GPU显存来调整。Teacher Forcing比率在训练序列生成模型时解码器当前步的输入可以是真实标签Teacher Forcing也可以是上一步自己的预测。使用一个比率如0.9来控制高比率有助于稳定训练初期但可能导致曝光偏差Exposure Bias。可以尝试在训练后期逐渐降低该比率。4.3 模型评估与指标解读训练不能只看训练损失必须在独立的验证集上评估。HMER常用的评估指标有表达式识别率 (Expression Recognition Rate, ExpRate)这是最核心的指标。它要求模型预测的整个LaTeX序列与真实标签完全一致才算正确。这非常严格一个空格错了都不行。词错误率 (Word Error Rate, WER)将LaTeX序列按词Token拆分后计算编辑距离插入、删除、替换的次数占真实词数的比例。它比ExpRate宽松一些能反映部分正确的预测。BLEU Score从机器翻译借鉴来的指标衡量预测序列与参考序列的n-gram重合度。在HMER中也有一定参考价值。在eval.py脚本中通常会实现这些指标的计算。我的经验是首要关注ExpRate因为它直接反映了模型能否完整准确地识别公式。WER可以作为辅助帮你分析模型常犯的错误类型是漏符号还是多符号。实操心得训练初期每1-2个epoch就在验证集上跑一次评估并保存验证集ExpRate最高的模型best_model.pth。同时将训练损失和验证损失绘制成曲线图这是诊断模型是否过拟合/欠拟合的最直观工具。如果验证损失很早就开始上升而训练损失持续下降那就是过拟合的典型信号。5. 实战排坑从环境配置到预测部署理论说得再多不如动手跑一遍。这一部分我结合自己复现时遇到的具体问题给出从零开始的排坑指南。5.1 环境配置与依赖安装首先确保你的Python环境推荐3.8和PyTorch版本建议1.9匹配。然后根据项目requirements.txt安装依赖。# 克隆代码仓库 git clone [CAN项目仓库地址] cd CAN-HMER # 创建并激活虚拟环境推荐 conda create -n can-hmer python3.8 conda activate can-hmer # 安装PyTorch请根据你的CUDA版本去官网选择命令 pip install torch torchvision torchaudio # 安装其他依赖 pip install -r requirements.txt常见坑点1CUDA版本不匹配。如果你用GPU训练务必确保安装的PyTorch版本支持你系统上的CUDA版本。用nvcc --version和python -c import torch; print(torch.version.cuda)来检查。常见坑点2缺少系统库。一些图像处理库如OpenCV或编译包可能需要系统级依赖。在Ubuntu上你可能需要sudo apt-get install libgl1-mesa-glx。5.2 训练过程中的典型问题与解决Loss为NaN或突然变得巨大原因梯度爆炸。这在RNN/LSTM中比较常见。解决在训练代码中已经提到的梯度裁剪clip_grad_norm_是关键。将裁剪阈值clip_grad设为一个小值如5.0或10.0。同时检查学习率是否过高可以尝试降低学习率。Loss下降很慢甚至不降原因学习率太小、模型初始化不当、数据预处理有问题如图像未归一化、标签词表构建错误。解决逐步调大学习率试试从1e-4到3e-4。检查数据加载流程打印几个样本看看图片张量是否在合理范围如归一化到[0,1]标签索引是否超出词表范围。检查词表确保训练标签中的所有符号都包含在词表中没有出现UNK未知符号索引。可以用一个极小的数据集如10张图先过一遍训练看是否能过拟合训练Loss快速降到接近0如果能说明模型和数据管道基本是通的。训练集Loss很低但验证集ExpRate几乎为0原因严重的过拟合或者训练集和验证集的数据分布差异极大。解决增加数据这是最根本的。收集更多数据或使用更激进的数据增强。正则化在模型中加入Dropout确保编码器和解码器的Dropout已启用或增加权重衰减Weight Decay。简化模型如果数据量很少尝试减小decoder_dim或使用更小的编码器如DenseNet-100。检查数据泄露确保训练集和验证集是严格分开的没有重复的图片。5.3 模型预测与可视化调试训练完成后使用predict.py脚本对单张图片进行预测。这个脚本会加载保存的最佳模型best_model.pth和词表进行前向传播并将输出的索引序列转换回LaTeX字符串。一个实用的调试技巧可视化注意力图。CAN模型的组合注意力机制是可解释的。你可以修改预测代码将解码过程中每一步的K个注意力权重图alpha_k保存下来并叠加到原图上。这能让你直观地看到模型在生成每个符号时到底“看”了图片的哪些部分。如果发现注意力图是散乱的或者没有聚焦到正确区域那可能意味着模型没有学到有效的空间关系需要回头检查数据或模型结构。# 在decoder的forward函数中返回注意力权重 def decode_step(self, ...): ... context, alphas self.attention(decoder_hidden, encoder_out) # alphas形状: [batch, K, height*width] ... return scores, decoder_hidden, alphas # 在预测时收集每一步的alphas all_alpha_maps [] for t in range(max_len): scores, hidden, alphas decoder.decode_step(...) all_alpha_maps.append(alphas) # 保存起来 # 预测结束后将all_alpha_maps转换为图像进行可视化5.4 用自己的数据训练检查清单为了确保流程顺畅这里提供一个自查清单[ ]数据图片已统一转为灰度图尺寸适中背景干净。[ ]标签已生成格式正确的标签文件如图片名\t LaTeXLaTeX序列的分词方式与代码要求一致。[ ]词表已使用训练集标签成功构建词表文件.pkl并确认无UNK问题。[ ]Dataset已修改__getitem__方法能正确读取自己的图片和标签并返回(image_tensor, target_indices)。[ ]配置已修改配置文件.yml更新了数据路径、词表路径、图片尺寸、批次大小等参数。[ ]环境PyTorch、CUDA、所有Python依赖已正确安装。[ ]试运行用小批量数据如2-4张图运行一个训练epoch确保没有报错如张量形状不匹配、索引越界并且损失可以计算。[ ]监控已准备好记录训练/验证损失和指标可以用TensorBoard或简单打印日志。[ ]备份已备份原始代码和配置文件开始在独立分支上进行修改。遵循以上步骤你应该能够顺利地将CAN模型应用到自己的手写公式数据集上。这个过程需要反复调试和耐心特别是数据准备和预处理阶段往往决定着模型性能的上限。当看到自己训练的模型第一次正确识别出一个复杂公式时那种成就感绝对是值得的。
返回列表