简介:基于ResNet与Transformer模型的手写数学公式识别项目源码,面向深度学习初学者、科研人员及课设/大作业开发者,针对手写公式图像中的符号定位与结构理解难题,采用ResNet提取视觉特征、Transformer自注意力机制建模符号排列关系,形成一套可运行、可扩展的识别方案。压缩包共40个文件,含19个Python脚本、pyc编译文件、yaml配置与说明文档,大小约4.21MB,代码按datamodule、model及训练/测试脚本分区,另有result.zip与附赠内容,便于模块化阅读和二次开发。目前已有152人浏览学习,项目为高分大作业并获导师认可,经严格调试、运行稳定,适合课设拓展或入门研究使用。读者可借助该源码快速上手上述两种模型的工程化组合,理解手写公式识别从数据预处理、特征编码到解码输出的完整流程,并可作为课程设计或论文实验的参考基线,也能为后续改进提供清晰起点。
1. 手写数学公式识别:为什么 ResNet + Transformer 是绕不开的组合
拿一份手写数学公式的图片丢给程序,要它输出一行 LaTeX 代码,比如得\frac{a}{b} + \sqrt{x}而不是一箩筐框出来的字符位置——这是手写数学公式识别(HMER)任务和普通 OCR 最大的分水岭。做过的人都知道,公式识别难点不在“认字”,在“结构”:上标下标、分式横线、根号嵌套,这些拓扑关系用纯 CNN 很难端到端建模,用纯序列模型又看不懂图像里像素级的细节。
这份基于 ResNet 与 Transformer 的 Python 项目,走的正是目前 HMER 里最主流的一条技术路线:ResNet 负责把图像压成粗粒度/细粒度的视觉特征,Transformer 解码器负责把特征序列翻译成 LaTeX 标记序列。它适合两类读者:一是课程设计或毕业设计需要快速跑通一个高完成度项目的学生;二是想在图像到序列任务里验证 Transformer 能力的工程师。下面从选型理由、数据预处理、推理部署到踩坑记录一步步拆开讲。
2. 选型拆解:ResNet 编码、Transformer 解码,这对组合为什么成立
2.1 残差卷积:公式图像里的字符到底怎么被“读”出来的
手写公式图像和自然场景图像最大的不同在于:字符小、密度高、结构符号(根号、括号、分式线)相互交叉。如果直接用 VGG 那种几十层纯卷积去提特征,网络加深后梯度消失问题会直接让训练崩掉。ResNet 的残差连接保证了每个 block 的输出至少包含输入的恒等映射,这让特征提取器在 50 层以上依然能稳定收敛。
在这类项目里,ResNet 通常不是拿来做分类尾巴的,而是用来当特征金字塔的地基。典型的做法是取 ResNet 的若干 stage 输出,形成不同分辨率的特征图——低层特征分辨率高、细节好(适合看小字符和连笔),高层特征语义强(适合看分式结构)。实际代码里,ResNet 最后一个 stage 输出的特征图会进一步压缩通道数,再拉平成序列喂给 Transformer。
import torch import torch.nn as nn from torchvision import models class ResNetEncoder(nn.Module): def __init__(self, d_model=512): super().__init__() resnet = models.resnet50(pretrained=True) # 去掉最后的分类头和池化,保留 conv1 -> layer4 self.conv1 = resnet.conv1 self.bn1 = resnet.bn1 self.relu = resnet.relu self.maxpool = resnet.maxpool self.layer1 = resnet.layer1 self.layer2 = resnet.layer2 self.layer3 = resnet.layer3 self.layer4 = resnet.layer4 # 把通道数统一映射到 d_model self.proj = nn.Conv2d(2048, d_model, 1) def forward(self, x): x = self.maxpool(self.relu(self.bn1(self.conv1(x)))) x = self.layer1(x) x = self.layer2(x) x = self.layer3(x) x = self.layer4(x) x = self.proj(x) # [B, d_model, H, W] b, d, h, w = x.shape x = x.flatten(2).permute(2, 0, 1) # [H*W, B, d_model] return x, (h, w)这里最关键的一步是最后的flatten(2).permute(2, 0, 1)。它把二维特征图展开成一维序列,得到[H*W, B, d_model],其中H*W就是 Transformer 解码器要处理的序列长度。对一张缩放到 224×224 的输入图,经过 ResNet 下采样后特征图通常是 7×7,展开后序列长度是 49,相当短,Transformer 跑起来非常快。
参数说明:d_model=512是 Transformer 内部的特征维度,也是 ResNet 输出通道统一映射到的宽度,这个值直接决定解码器参数规模。pretrained=True表示使用 ImageNet 上预训练好的权重做初始化,工程上强烈建议打开——公式字符虽然和 ImageNet 类别不同,但底层边缘、纹理、笔画的滤波器是可以迁移的。
细看网络结构会发现layer4输出的分辨率是最低的,只有输入的 1/32。如果公式图里字母格外小,这么粗的特征会丢笔画细节。现实中有的项目只取到layer3,把下采样倍数降到 1/16,用分辨率换感受野。这个取舍没有绝对标准,我的做法是:图片短边小于 1000 像素就保留 layer4,小于 500 像素就砍到 layer3,保证特征图上每个格子至少对应原图 16×16 区域。
2.2 自注意力解码:为什么能对齐图像特征和 LaTeX 序列
ResNet 生成特征序列后,Transformer 解码器负责逐步生成 LaTeX 标记,它不再依赖固定窗口的卷积,而是对已经生成的所有历史标记和视觉特征做全局注意力。这一步对公式识别尤其重要:预测\frac的时候,解码器需要看到分数线的位置;预测}的时候,需要回看之前出现过的{。
解码器的输入有两个:一是视觉特征序列(作为交叉注意力的 Key/Value),二是已经生成标记的嵌入向量(作为自注意力的 Query)。训练阶段用真实序列做 Teacher Forcing,推理阶段用上一步的预测结果作为下一步输入。
import math import torch.nn as nn class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len=500): super().__init__() pe = torch.zeros(max_len, d_model) position = torch.arange(0, max_len).unsqueeze(1).float() div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)) pe[:, 0::2] = torch.sin(position * div_term) pe[:, 1::2] = torch.cos(position * div_term) self.register_buffer('pe', pe) def forward(self, x): return x + self.pe[:x.size(0)]位置编码在这份资源里的地位被很多人低估。Transformer 结构本身不含顺序信息,特征图拉成的序列和 LaTeX 标记序列都是有序的,缺少位置编码会导致分式分子分母顺序错乱。上面这段用的是标准正弦位置编码,对于公式这种序列长度通常小于 200 的场合足够用。
而针对二维码特征图展开后的序列,原论文在实践中还有一个变体做法:把特征图的横纵坐标直接当作额外 token 加入注意力计算,相当于给 Transformer 提供二维位置先验。常见做法是在ResNetEncoder输出的[B, d_model, H, W]特征上拼接一个坐标卷积层,生成横坐标图和纵坐标图各一通道,然后 concat 进特征。视觉特征的位置编码和文本位置编码本质上解决同一个问题,但前者更细粒度——公式字符的空间关系远比纯文本句子紧密。
2.3 粗粒度与细粒度特征:公式识别的核心拉锯战
读到这里会碰到这个概念:手写公式解码 LaTeX 时,粗粒度特征负责定结构骨架,细粒度特征负责认字符笔画。只给解码器最后一层特征,模型对\sum和\int这种外形接近的大符号容易混淆;只给第一层特征,模型对整体结构就失去全局感。
要兼容两者,常见的处理方式是融合多尺度特征后再进解码器。不是简单相加,而是把不同 stage 的特征上采样/下采样到统一分辨率后做通道拼接。我在复现这个项目时习惯把 layer2(细粒度)和 layer4(粗粒度)都取出,layer4 的语义信息通过双线性插值放大到与 layer2 相同尺寸,然后沿通道拼接,再投影到d_model。
这个设计的收益很直观:细粒度特征让模型准确地观察到分数线上下缘的像素落差,粗粒度特征让模型不在根号边界上犯方向性错误。代价是序列长度会增加,注意力计算量跟着上涨。如果显存吃紧,优先砍掉细粒度分支,保留粗粒度分支,结构稳定性优先于字符细节。
3. 搭建一个能用的训练管线:数据、模型、损失函数怎么串起来
3.1 数据组织:图像-序列配对,缺一不可
手写公式识别的训练数据本质上是一个配对集合:每张图片对应一行 LaTeX 标注。项目里这份数据通常来自公开数据集或自行采集。完整管线的第一步是把图片路径和 LaTeX 标签做成索引文件。
import os import pandas as pd # 假设图片在 images/ 目录,labels.csv 中每行是 "图片名, LaTeX标注" df = pd.read_csv('labels.csv', header=None, names=['img_name', 'latex']) pairs = [] img_dir = 'images' for _, row in df.iterrows(): path = os.path.join(img_dir, row['img_name'] + '.png') if os.path.exists(path): pairs.append((path, row['latex'])) train = pairs[:int(len(pairs)*0.8)] val = pairs[int(len(pairs)*0.8):] print('train:', len(train), 'val:', len(val))这段代码做的事是把标注文件和实际图片文件做一次存在性校验,然后按 8:2 切分训练集与验证集。别小看这个校验步骤,公式数据集经常出现标注里写了图但文件丢失的情况,提前过滤可以避免训练中途 FileNotFoundError 打断流程。
3.2 预处理流水线:统一尺寸和归一化
LaTeX 标签也要做 tokenization。需要把它拆成 token 序列,因为模型不是按字符预测的,而在 LaTeX 语法里\frac是一个整体 token,不能拆成\f、\r。数据预处理模块通常要维护一个词汇表:从训练集所有 LaTeX 标签里统计 token 频次,过滤低频项,然后建立 token 到整数索引的映射。
from transformers import PreTrainedTokenizer # 常见做法:用简单的正则按 LaTeX 语法切分 import re def latex_tokenize(latex_str): # 将 \frac 等命令视为整体,括号单独成 token tokens = re.findall(r'\\[a-zA-Z]+|[{}]|.', latex_str) return tokens sample = r'\frac{a}{b}' print(latex_tokenize(sample)) # 输出示例: ['\\frac', '{', 'a', '}', '{', 'b', '}']\\[a-zA-Z]+匹配所有以反斜杠开头的 LaTeX 命令;[{}]把大括号单独提出来;.兜底匹配单个字符。这样\frac不会碎成\+f+r+a+c。需要特别注意的是,正则中\\表示匹配一个真正的反斜杠,所以\\[a-zA-Z]+写法没问题,但不要写成\[a-zA-Z]+,后者会变成匹配方括号。
3.3 训练循环:Teacher Forcing 与损失遮蔽
训练时用 Teacher Forcing——把真实标签序列的一部分输入解码器,让模型预测下一个 token。计算损失时只关注 LaTeX 有效 token,<pad>位置必须遮蔽掉,否则模型会学到大量无效内容拉低指标。
import torch import torch.nn as nn def train_one_epoch(model, loader, optimizer, criterion, pad_idx, device): model.train() total_loss = 0 for images, latex_seq in loader: images = images.to(device) latex_seq = latex_seq.to(device) # 准备解码器输入:去掉最后一个 token input_seq = latex_seq[:, :-1] # 准备监督目标:去掉第一个 token target_seq = latex_seq[:, 1:] optimizer.zero_grad() # logits: [batch, seq_len, vocab_size] logits = model(images, input_seq) loss = criterion( logits.permute(0, 2, 1), target_seq ) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0) optimizer.step() total_loss += loss.item() return total_loss / len(loader)input_seq和target_seq错位一位,是 Teacher Forcing 的标准节奏,模型永远是在看过前面真实 token 后预测下一个。clip_grad_norm_是对梯度做裁剪,公式识别任务里 LaTeX 序列较长,梯度爆炸是家常便饭,不裁剪的话 loss 会在某一轮后突然变成 NaN。
损失函数用CrossEntropyLoss时,注意ignore_index=pad_idx要传进去,不传的话,<pad>位置的错乱预测也会计入 loss,产生梯度噪声。
4. 数据集与预处理:最容易翻车的预处理细节
4.1 图像尺寸与字符密度的矛盾
公式图片的宽高比极不均匀:有的是一行短公式,有的是多行复杂分数。如果强行 resize 到正方形,字符会被压扁拉长,识别率骤降。常见做法是先等比例缩放,让长边等于 224,短边不足的地方用白色像素填充到目标尺寸,再训练。也可以不做 pad,直接在 dataloader 里用 batch sampler 把尺寸接近的图分到同一 batch,减少计算浪费。
手写公式字符相对密集,图像增强里的随机裁剪要克制。随机旋转超过 5 度会毁掉上下标关系,随机颜色扰动基本没用,因为公式图通常是白底黑字。我一般只用两种增强:轻微仿射变换和随机亮度扰动,幅度都控制在 3% 以内,目的是模拟不同书写工具的颜色差异,而不是制造新的变形。
4.2 LaTeX 标注的清洗策略
公开数据集的 LaTeX 标签常有噪声:空格位置不对、宏包命令不一致、成对括号缺失。清洗时统一做四件事:去掉多余空格,兼容\dfrac和\frac,把\left(与\right)标准化为普通括号,过滤掉数据集中出现次数少于 3 次的罕见 token。前三个提升一致性,第四个防止词汇表膨胀到几千导致训练不收敛。
4.3 数据加载的性能瓶颈
公式识别图片数量大,每张图又要先做 resize、再归一化、再增强,CPU 处理速度很容易跟不上 GPU。PyTorch 里DataLoader的num_workers要设到 4 以上,同时开启persistent_workers=True,避免每轮都重新创建子进程。如果机器内存够,可以加一个轻量缓存,把加载过的处理结果以字典形式存放,虽然多吃内存,但训练时间能缩短三分之一。
5. 避坑:五个真实踩过的坑,现象、原因与解决
5.1 特征图展开后序列方向不对导致结构错乱
现象:模型预测的 LaTeX 结构整体反了,分子分母颠倒,根号内部内容跑到外面。
原因:flatten(2)是按行优先展开的,也就是先走完第一行再走第二行。如果你的特征图是[H, W],展开后序列中位置i对应原图坐标(i // W, i % W)。但 Transformer 解码器默认认为序列从左到右、从上到下是自然顺序,如果代码里先flatten后permute时维度搞混,特征图可能被转置了。
解决:把 ResNet 输出的特征图按[B, d_model, H, W]显式检查一下打印出的形状。然后用torch.arange手动生成坐标序列,确认位置0对应特征图左上角,位置W-1对应右上角。写个小测试:把一张只有左上角有黑点的二值图送入编码器,查看加权热力图的重心坐标是否也落在左上角。
5.2 梯度裁剪不当导致 loss 变成 NaN
现象:训练到第 20 轮附近,loss 突然变成 NaN,重启后又能跑几轮再次崩掉。
原因:clip_grad_norm_裁剪的是梯度的二范数,但如果某一层的梯度本身已经包含了 NaN,裁剪是针对 NaN 之外的数值做的,NaN 会绕过裁剪继续传播。这类情况通常发生在反向传播时某一步数值溢出,多见于 logits 过大时 softmax 求幂溢出。
解决:先把max_norm从 5.0 降到 1.0,降低梯度幅值,看是否还崩。同时检查学习率是否超过 1e-4,Transformer 对这种高维序列任务比 CNN 敏感得多。在损失计算后加一行torch.isnan(loss)检测,如果为真就跳过这一步更新并打印当前 step,比盲目调参快得多。
5.3 位置编码长度不够导致推理阶段直接崩
现象:训练时一切正常,推理时遇到长公式,报索引越界错误。
原因:训练阶段 LaTeX 序列长度被max_len限制住了,比如设定为 200。但推理时公式真实序列可能超过 200,位置编码矩阵只有 200 行,越界就是必然。
解决:训练阶段把max_len设为 600,给足余量,毕竟推理时的序列长度完全由输出决定。另外在推理循环里加判断,如果预测长度达到位置编码上限就停止生成,而不是让它报错。
5.4 预训练 ResNet 的 BatchNorm 统计量漂移
现象:加载预训练 ResNet 后,训练初期 loss 不降反升,验证集准确率几乎为 0,且模型收敛极慢。
原因:预训练统计量来自 ImageNet 数据分布,公式图片与自然图像差别大,BatchNorm 层的 running_mean 和 running_var 需要较长时间适应新分布。更隐蔽的是,有些代码把 resnet 设置为requires_grad=False,只训练后面的 Transformer,这样 BatchNorm 统计量不更新,特征分布始终偏向自然图像域,Transformer 学到的是和公式域不一致的特征。
解决:要么在加载预训练权重后把所有nn.BatchNorm2d的track_running_stats保持默认开启并训练 2 个 epoch 做 warmup,让running_mean先适应新分布;要么干脆不用预训练权重,从头训练 ResNet,代价是收敛慢,但最后效果通常差不多,因为公式图的视觉分布和 ImageNet 差异实在太大。
5.5 数据集中 LaTeX 语法本身不规范
现象:验证集 loss 很低,但渲染出来的结果图语法错误一片红,\frac缺参数导致 PDF 编译失败。
原因:标注数据里有的\frac后只跟了一个{...},有一半的标签真值这种情况下模型被训练成“预测不完整命令”的习惯,解码时输出的 LaTeX 自然不合法。
解决:数据清洗时用正则做语法规则检查,扫描不闭合的\frac{、\sqrt{和多余括号。更粗暴的方案是训练结束后用渲染库(比如 matplotlib 的 mathtext)对预测结果做语法验证,渲染不通过的按错误处理。这个方法不需要额外的标注数据,是工程上最快找到问题的手段。
6. 把结果变稳:验证模型可信度的三个硬办法
跑通了训练和推理,只能说模型“能出声”,不能说它“靠谱”。衡量手写公式识别任务,准确率并不是数字对就是全部,得看渲染出来的 LaTeX 能不能被解析器认出来。以下三个验证习惯是我在复现该项目时坚持下来的。
第一个验证办法是渲染回读。把预测出的 LaTeX 字符串交给matplotlib渲染成图片,再与输入的图片做尺寸比对。如果渲染出来的图比例和原图相差过大,说明结构识别有偏差;如果渲染直接报错,说明 token 序列语法有错误。这样做的好处是把一个“猜对不对”的问题变成一个可自动检验的对错。
import matplotlib.pyplot as plt def render_latex(latex_str, output_path='render.png'): fig = plt.figure(figsize=(6, 2)) t = fig.text(0.5, 0.5, f'${latex_str}$', horizontalalignment='center', verticalalignment='center', fontsize=20) try: fig.canvas.draw() fig.savefig(output_path, dpi=100, bbox_inches='tight') plt.close(fig) return True except Exception as e: plt.close(fig) return Falsefig.text的字符串要包在$里,matplotlib 才把它当公式渲染。如果latex_str本身语法不合法,draw阶段会抛出异常,捕获后返回False即可。值得注意的是,matplotlib 的 mathtext 支持的 LaTeX 子集比完整 LaTeX 编译器小,有些真值标注通过完整 LaTeX 编译没问题但在 mathtext 会炸,所以做验证时不要一棒子打死,渲染失败只能说明大概率有问题,不能直接判错。
第二个验证办法是混淆式测试。把一张测试公式图做 10 度旋转,看输出是否还是同样的公式。如果旋转后识别结果大改,说明模型过度依赖方向特征而不是字符特征。更有效的是把图片做垂直翻转,公式识别模型如果输出不变就说明它根本没有学习到公式的结构逻辑。这个测试不花多少时间,却能在正式提交项目前暴露模型是记住了训练集还是真的学到了拓扑关系。
第三个验证办法是全流程推理测试脚本,把所有环节串起来输出一份报告,包含输入图片路径、预测 LaTeX、渲染状态、耗时。这里耗时数据要重视:CPU 推理和 GPU 推理差距极大,如果资源里自带推理脚本,先在 CPU 上验证一次功能完整性,再切 GPU。
最终我习惯在每次训练结束后,强制跑一遍渲染回读脚本,把验证集里所有预测结果渲染成图片,挑出渲染报错的样本做案例分析。从那次遇到\frac缺参数导致大批结果无法渲染开始,我再没只盯着 loss 下降就判断模型合格,而是把“预测结果能不能渲染成合法文档”列进了验收标准。希望这个方法也能帮你的模型诊断省点事。
本文还有配套的精品资源,点击获取