ARTICLE DETAIL

资讯详情

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

手写数学公式识别实战:ResNet+Transformer架构与工程落地

手写数学公式识别实战:ResNet+Transformer架构与工程落地 简介基于ResNet与Transformer架构的手写数学公式识别Python源码项目面向高校学生在深度学习课程大作业、毕业设计或算法竞赛中的需求实现从手写图像到LaTeX公式序列的端到端识别。压缩包共32个文件以19个Python源码为核心按数据模块、模型模块、训练与测试模块组织覆盖数据加载、词典构建、编码器、解码器、位置编码、批处理训练及验证推理等完整流程另含txt字典与说明、yaml/cfg配置以及pyc缓存文件整体仅87KB轻量易部署。该项目为个人高分大作业经导师指导认可并完成严格调试确保可直接运行。目前已有646人学习下载通过学习可拆解ResNet特征提取与Transformer序列解码的协作机制掌握公式识别任务中字典生成、训练脚本、验证逻辑等工程细节。目录结构清晰既适合作为课程设计参考与算法复现实验也可作为后续优化与扩展的基础框架。1. 手写数学公式识别为什么是 ResNet Transformer而不是纯 CNN手写数学公式识别这个方向本质上不是“识别字符”而是“把图片翻译成结构化的 LaTeX 序列”。标题里把 ResNet 和 Transformer 放在一起正是因为这条路线在工程上最成熟ResNet 负责把笔画、根号、上下标这类视觉线索变成特征图Transformer 解码器负责生成\frac、^、_这种带嵌套结构的 token 序列。你拿纯 CNN 做分类只能得到符号列表却得不到它们之间的层级关系拿纯 Transformer 直接吃像素又太浪费算力。这套 Python 源码要解决的问题就是让一张手写公式图片经过网络后直接输出可供 LaTeX 编译器渲染的文本。这个项目适合三类人准备在 CROHME 数据集上复现指标的研究生想把试卷扫描自动录入做成网页 API 的 OCR 工程师以及在评估“公式识别到底能不能落地”的产品负责人。它不只是一个模型脚本而是把数据准备、训练、推理串起来的完整工程。下面我从架构选型开始逐步拆开这套识别方案的做法、参数和坑。2. 架构选型把 ResNet 特征图变成 Transformer 能吃的序列标题里的 ResNet 不是用来做分类的它在这里承担“视觉编码器”的角色。Transformer 也不是用来做分类的它承担“序列解码器”的角色。两者之间的接缝设计决定了模型能不能学会公式的结构。2.1 编码器与解码器的接口怎么定常见做法是让公式图片先经过 ResNet拿到一个形如[B, C, H, W]的特征图再用卷积把通道数压缩到 Transformer 需要的d_model然后把特征图在空间维度上展开成[H*W, B, d_model]的序列。这个序列就是 Transformer 解码器的 memory。这里有第一个关键选择用 ResNet 的哪一层输出。只看最后一层感受野大能抓住整个公式的布局但笔画的细节已经在多次下采样中丢得差不多了用较浅的层细节保留好却缺少全局上下文。许多实现会在 ResNet 后面接一个 1x1 卷积把特征图通道从 2048 压到 256同时保持空间尺寸不变。这样 Transformer 的注意力矩阵大小只由H*W决定训练速度也能接受。我一般会在把这套流程接通之前先打印一次各个阶段的特征图尺寸。假设输入图片是64x512ResNet50 的 C5 输出 stride 是 32那你会得到2x16的空间尺寸展开后只有 32 个位置这对注意力来说是很轻的。但如果输入图片不缩放直接送 512x512stride 16 的 C4 特征图就是32x321024个位置解码器每生成一个 token 都要和 1024 个位置做交叉注意力训练时间和显存会明显上升。2.2 粗粒度特征与细粒度特征怎么取舍手写公式识别对特征粒度的要求比普通 OCR 更苛刻。下标小到只有几个像素根号边界又往往横跨整个公式单层特征很难同时兼顾这两种跨度。粗细粒度特征的融合是这类项目里最容易拉开差距的部分。单一使用高层特征时常见问题是\sqrt{}的右边界被模型忽略上标被当成普通字符单一使用低层特征时整体结构又会乱掉。比较实用的做法是把 ResNet 的 C3、C4、C5 三层输出做一个轻量融合类似 FPN 的思路。对每层输出做 1x1 卷积统一通道再上采样到同一尺寸逐元素相加。这样做不会显著增加参数量但对根号内嵌套公式、上下标对齐这些细节的提升非常明显。如果源码里没有 FPN 融合也可以用更简单的方式把 C4 和 C5 的输出在通道维度拼接再用 1x1 卷积把通道数压回d_model。这样的效果比单层 C5 好很多改动却只有几行。2.3 位置编码一维编码为何不适合二维公式Transformer 本身没有顺序概念所以必须给每个特征位置加位置信息。标准 Transformer 用的是正弦一维位置编码但公式是一个二维结构单维编码无法区分“同一行靠右的位置”和“下一行的同一水平位置”。尤其在处理\frac{a}{b}这类上下结构时一维位置编码很容易把分子和分母的位置信息搞混。一个可靠的做法是使用二维可学习位置编码分别对应特征图的行坐标和列坐标。下面这段代码是这类工程里常见的实现import torch import torch.nn as nn class PositionalEncoding2D(nn.Module): def __init__(self, d_model, max_height8, max_width32): super().__init__() self.row_emb nn.Parameter(torch.randn(max_height, d_model // 2)) self.col_emb nn.Parameter(torch.randn(max_width, d_model // 2)) def forward(self, x, h, w): # x 的形状: (B, H*W, d_model) row self.row_emb[:h].unsqueeze(1) # (h, 1, d_model//2) col self.col_emb[:w].unsqueeze(0) # (1, w, d_model//2) pos torch.cat([ row.expand(h, w, -1), col.expand(h, w, -1) ], dim-1).reshape(h * w, -1) return x pos.unsqueeze(0)这段代码把d_model256拆成两半前 128 维编码行坐标后 128 维编码列坐标然后广播到整个特征图。这里的max_height和max_width不需要设得太大能够让程序覆盖训练时的最大特征图尺寸即可。如果设置过大可学习参数会增加但真正被用到的只有前面一部分。需要注意这个位置编码是加在编码器输出上的。Transformer 解码器在交叉注意力里会隐式学习“哪个特征位置对应哪部分公式结构”所以位置编码加得对不对直接影响上下标的对齐效果。调试时可以只看编码器输出的热图如果同一行公式的特征位置被映射到相近向量那位置编码基本没问题。2.4 LaTeX 规范化决定训练能不能收敛的细节公式识别的输出不是字符序列而是 LaTeX token 序列。同一个公式可能有多种等价写法比如a^2和a^{2}在视觉上完全一样。训练数据里如果混用这两种写法模型会无所适从。源码工程里通常会有一个专门的 tokenizer把 LaTeX 字符串拆成 token并统一成一套规范括号全部使用{}显式分组根号使用\sqrt{}或\sqrt[]{}指数统一写成^{...}。这样做是为了给模型一个确定性的目标序列。你只看损失曲线时可能看不出问题但验证集准确率会卡在某个值上不去这就是目标空间不干净造成的。另一点是 token 词典的控制。常见项目的词典一般控制在 200 到 500 个 token 之间包含常见字母、数字、运算符、希腊字母和控制序列。词典太大时解码器最后一层分类头的计算会成为瓶颈所以不要盲目把全部 Unicode 字符都放进去。3. 训练工程落地数据、参数、训练循环架构定下来之后下一步是把训练工程跑通。这个标题里的 Python 源码如果只是一个模型定义其实价值不大真正有价值的是里面的数据准备脚本、训练配置和推理代码。3.1 数据准备CROHME 与自建数据集的目录规范公开数据集里最常用的是 CROHME里面的手写公式有 LaTeX 标注但原始格式是 InkML需要先渲染成图片。比较标准的目录结构是这样data/ images/ train/001.png val/002.png labels/ train.json val.jsontrain.json里一般保存图片路径和对应的 LaTeX 字符串。读取时先把图片缩放成固定高度宽度按比例变化再填充到固定宽度。高度通常取 64宽度取 256 到 512 之间。高度太小会丢失小符号细节高度太大则序列太长影响效率。如果你是自己采集数据建议保留写入时的墨迹轨迹而不是直接存成品图片。采集后做一个预处理把轨迹渲染成白底黑字的图片再按非零边界裁剪统一缩放到目标尺寸。把轨迹文件本身也保存下来以后做数据增强时可以直接对轨迹做仿射变换再渲染成图效果比增强图片更接近真实手写。3.2 关键训练参数学习率、warmup、梯度裁剪与输出长度Transformer 系列模型的训练参数和纯 CNN 差别很大。直接用Adam配上固定学习率通常会出现前期发散、后期不收敛的情况。我更推荐下面这种配置# config.py CFG { img_height: 64, img_width: 512, max_len: 128, d_model: 256, enc_layers: 6, dec_layers: 6, nhead: 8, resnet_name: resnet50, pretrained: True, dropout: 0.1, batch_size: 16, lr: 1e-4, warmup_steps: 1000, clip_grad_norm: 5.0, beam_size: 5, use_amp: True, }这里的d_model256是 Transformer 内部的特征维度ResNet 输出的 2048 维特征需要先被压到这个维度再送入解码器。nhead8意味着 256 能被 8 整除注意力头数为 8。warmup_steps是预热步数前 1000 步学习率从很小的值线性涨到lr之后再按步数倒数衰减。这个机制对 Transformer 非常关键尤其是从零训练而不是加载预训练解码器的时候。max_len控制解码器最多生成的 token 数。CROHME 里大部分公式在 80 个 token 以内设成 128 足够如果你的数据里有特别长的多项式推导建议设到 256。但要明白max_len越大训练时的注意力矩阵也越大显存占用是平方级增长的。3.3 训练循环与评估口径训练循环的核心是带 teacher forcing 的交叉熵损失解码器每一步的输入是真实的前一个 token预测目标是下一个 token。写成代码大致是这样def train_step(model, optimizer, x, y, pad_idx, scaler): model.train() optimizer.zero_grad() with torch.cuda.amp.autocast(): # 解码器输入是 y[:, :-1]预测目标是 y[:, 1:] logits model(x, y[:, :-1]) loss nn.functional.cross_entropy( logits.reshape(-1, logits.size(-1)), y[:, 1:].reshape(-1), ignore_indexpad_idx, ) scaler.scale(loss).backward() nn.utils.clip_grad_norm_(model.parameters(), 5.0) scaler.step(optimizer) scaler.update() return loss.item()ignore_indexpad_idx很关键。公式序列长短不一一个 batch 内部会把短句子填充到等长填充位置不参与损失计算否则模型的大量梯度会来自无意义的padtoken。clip_grad_norm设为 5.0 是为了防止长序列上梯度爆炸公式识别里梯度的绝对值往往比图像分类大很多不裁剪的话训练到中后期会突然出现 loss 变成 NaN。评估时不要只看 loss。每训练 1000 步挑几批验证集图片用 beam search 解出 LaTeX 字符串再算整句准确率或公式级准确率。这里有一个容易被新手忽略的点验证时的输入需要通过真实分词器解码成最终字符串而且没有标准化的话x^2和x^{2}会互相扣分。建议在评估脚本里做一步 LaTeX 再规范化把所有^{...}格式统一。4. 避坑记录手写公式识别最容易踩进的 5 个坑这一类项目我在实际跑的过程中踩过不少坑而且很多是现象相似、原因完全不同只能把现象、原因、解决路径一起说清楚。4.1 现象损失持续下降但整句识别率纹丝不动原因很容易被掩盖交叉熵损失是 token 级别的只要模型能预测对大部分简单字符损失就会好看但公式结构里的一个\frac出错后面整个分子分母都会错位整句准确率仍然为零。我见过不少项目止步在 30% 左右的整句准确率就是只盯 loss 不看解码结果。解决方法是把验证指标从“per-token accuracy”换成“整句匹配率”并且每几百个 step 就把模型输出的 LaTeX 渲染出来人工看几张。还有一个更有效的做法在训练后期用小学习率微调并增大 label smoothing 到 0.1让模型不要把概率过度集中到某些高频 token 上。这样做整句准确率通常能再涨两三个点。4.2 现象图片空白区域多模型乱输出符号手写公式图片往往四周有大片空白。ResNet 的特征图会把这些空白区域也编码成位置向量模型在交叉注意力时可以选择“看空白”导致输出的 LaTeX 里夹杂奇怪的起止位置。原因就是预处理阶段没有做白边裁剪或者裁剪后没有统一缩放到合适的尺寸。解决方式是在送入网络前按行投影和列投影找到非零像素的边界做一次 tight crop再填充到目标尺寸。这样模型看到的有效区域更大注意力也会更集中。建议在训练和推理时使用完全相同的裁剪逻辑否则训练时学到的“空白应对方式”和推理时不匹配。4.3 现象显存占用高得离谱训练直接 OOM手写公式图片宽度往往很大如果直接缩放成256x256的正方形ResNet 输出的特征图长度会非常可观。Transformer 解码器的自注意力是序列长度的平方复杂度图片越宽序列越长显存增长远超线性。解决方式是调整输入口径固定高度为 64宽度按比例缩放后再填充到 512。这样特征图只有2x1632个位置解码器序列长度也受max_len128限制整体显存占用可控。还可以开混合精度训练在多数显卡上能把显存占用压到原来的 60% 左右。若仍然不够减小 batch size同时用梯度累积来保持有效 batch 大小。4.4 现象长公式输出被截断解码陷入死循环推理阶段如果 beam search 在达到max_len时都没遇到eos程序会强制截断导致公式后半段丢失。更隐蔽的问题是模型反复输出同一个 token比如连续十几个空格或连续多个\left。原因有两个一是训练数据里长公式占比太少模型没见过足够多的长序列二是 beam search 的长度归一化处理不对。解决方式是在训练时不要对长度做过多的过滤保留真实的长度分布推理时给长度惩罚设置合适值。例如输出概率按每新增 token 的score / len^alpha来比较alpha取 0.6 到 1.0。如果发现模型倾向于输出非常短的公式就降低alpha如果输出过长且重复就提高alpha。4.5 现象训练时 teacher forcing 效果很好一换成自回归就崩这是一个典型的训练与推理不一致问题。训练时解码器每一步都输入真实 token模型可以偷懒只学会在正确答案条件下预测下一个 token推理时它要把自己生成的 token 再作为输入一旦某一步出错错误会向后传播输出变得不可控。解决这个方法最常见的是 scheduled sampling在训练过程中以一定概率把真实 token 替换为模型自己的采样结果。概率从 0 开始随着训练进度逐步上升到 0.1 或 0.15。注意这个概率不能太大否则训练不稳定。我一般会让它在后 30% 的训练轮次里线性增长并在每轮结束后用验证集做一次自回归解码观察稳定性变化。5. 部署与推理接口让 Python 源码真正能被调用训练完的模型最终要交给业务去调用不会有人真的去操作命令行跑预测。这一章说清楚模型导出、beam search 和 API 接口三件事。5.1 模型导出时的边界条件常见的做法是用torch.save保存完整的 weight 文件然后在服务端重新加载网络结构。文件不大复现也容易。但如果想用 TorchScript 或 ONNX 进一步优化就需要处理解码器中的循环调用。注意beam search 里有大量 while 循环、条件分支和集合操作这些逻辑用 TorchScript 追踪会有困难。更稳妥的做法是只导出编码器和“单步解码”部分编码器把图片变成 memory解码器根据当前 token 序列输出下一个 token 的 logitsbeam search 的循环留在 Python 层控制。这样可以被粗粒度特征注意力和细粒度特征网络共同优化模型在边缘设备上也能保持可控的内存占用。如果目标设备只支持 ONNX需要特别注意动态输入长度。公式图片的宽度和输出序列长度是动态的导出时要显式指定动态轴batch、height、width、sequence。否则 ONNX 会把分辨率固定死换一张不同宽度的图片就会报错。5.2 推理函数与 beam search 参数服务端的核心函数不复杂但 beam search 的细节决定最终效果。下面是一个可以用在工程里的最小实现def decode_beam(model, image, tokenizer, beam_size5, max_len128, alpha0.8): model.eval() with torch.no_grad(): memory model.encode(image.unsqueeze(0).to(device)) beams [(tokenizer.bos, 0.0)] # (token序列, 累计对数概率) for _ in range(max_len): candidates [] for seq, score in beams: if seq[-1] tokenizer.eos: candidates.append((seq, score)) continue logits model.decode_one_step(memory, torch.tensor([seq], devicedevice)) log_probs logits.log_softmax(dim-1)[0, -1] top_k log_probs.topk(beam_size) for token_idx, token_lp in zip(top_k.indices, top_k.values): candidates.append((seq [token_idx], score token_lp.item())) # 按长度归一化后的分数排序 candidates.sort(keylambda x: x[1] / (len(x[0]) ** alpha), reverseTrue) beams candidates[:beam_size] best beams[0][0] return tokenizer.decode(best)这里最重要的参数是alpha它控制长度惩罚的强度。alpha太小模型倾向于生成短公式因为它累积了较少的负对数概率alpha太大模型会产生冗长的重复内容。我一般先在验证集上网格搜索0.6、0.8、1.0三个值挑整句准确率最高的那个。beam_size也不是越大越好。beam size 从 1 提到 5 通常能带来 5 到 10 个百分点的整句准确率提升但再往上提升就变得很小而解码耗时接近翻倍。如果服务端的计算资源有限先用 beam size 5再用定点化或批量预测来补偿速度。5.3 JSON 接口中的图像约定给外部的识别接口一般用 JSON 格式请求里放一张 base64 编码的图片响应里返回识别出的 LaTeX 字符串和置信度。这里最容易出问题的是图像格式约定不一致。建议在接口文档里明确说明三点第一图片背景必须是白底黑字如果上游传来的是手机拍照的灰底图片先做二值化和反色第二公式内容不要在图片里带边框、水印或其他文字否则识别结果会出现多余字符第三请求图片的宽度上限最好由服务端预处理再裁剪而不是直接由客户端控制因为不同客户端生成的图片尺寸差异很大。置信度不要直接使用 softmax 概率的均值因为解码序列很长概率乘积会非常小。更实用的定义是“length-normalized score”即 beam search 里最终候选的累计对数概率除以长度然后映射到 0 到 1 区间。这个分数虽然不能和准确率一一对应但至少能用来做阈值过滤把低质量图片识别结果标记出来交给人工确认。6. 一个进阶技巧用 FPN 特征与二维位置编码提升低质量手写识别如果你的模型已经在标准数据集上跑通但一到真实手写场景就出问题我建议把注意力放在“细粒度特征”和“位置编码”的配合上。手写公式和印刷体最大的不同是笔画的粗细、斜度、断裂都不可控。根号内部的被开方式可能被写得几乎贴近根号线上标也可能和主体字符连成一团。单靠 ResNet 最后一层输出这种视觉歧义很难被消解。我习惯在工程里把 FPN 融合层的输出直接作为 Transformer 编码器的输入同时给每一层特征图都加上二维位置编码。低层特征保证了笔画边界清晰高层特征保证了结构关系正确二维位置编码则让解码器知道“这个特征来自图片的哪一行哪一列”。这个方法在验证集上不一定每次都涨分但对长根号、多层分数这类结构复杂样本的帮助很明显。另一个小技巧是关注注意力可视化。把解码器某一层交叉注意力的权重映射回原图你会直观看到模型在生成\frac时是否把注意力放在了分子和分母区域。如果注意力非常分散优先怀疑位置编码没有生效或者输入图片没有做自适应对比度增强。针对低质量图片可以加一个轻量预处理先做自适应阈值二值化再对笔画做一次膨胀让细小的手写笔迹连起来。千万不要在训练和推理时使用不同的预处理方式这是很多项目“训练正常、一到线上就翻车”的最大原因。我在做这类项目时最后留下一个习惯每次训练前先随机取出 20 张训练图片手动渲染出预处理后的结果挂在一个文件夹里训练过程中随时翻看。这样能发现很多模型之外的问题比如标注错误、图片拉伸变形、白边没有裁剪干净。预处理正确模型才可能正确。如果你正为手写公式识别的准确率卡在瓶颈而苦恼不妨先检查这两个地方位置编码是不是真的用上了二维信息特征提取是否太依赖 ResNet 的最后一层。把这两处补上再配合 beam search 的长度惩罚调优识别效果通常会有明显提升。希望帮到你。本文还有配套的精品资源点击获取
返回列表