简介:一份基于ResNet与Transformer架构的手写数学公式识别Python源码,属于高分课程大作业项目,适合深度学习、计算机视觉方向的开发者参考。项目已通过导师指导与验收,代码经过严格调试,可直接运行,覆盖数据加载与预处理、模型搭建、训练验证及公式识别预测的完整流程。压缩包共32个文件,以Python脚本为主,含19个py源文件、8个pyc编译文件与2个txt字典/结果文件,另附cfg和yaml配置文件,整体仅87KB,便于快速浏览和部署。代码按datamodule、model等功能模块组织,内含编码器、解码器、位置编码、配置加载与验证脚本,并附带单条公式识别结果文本,可帮助理解识别管线。目前已有646人学习下载,适合想跑通手写公式识别基线、在现有结构上做模块替换或二次开发的读者使用。
1. 手写公式识别项目到手:先看清 ResNet+Transformer 在解决哪件事
手写数学公式识别,是把一张手写公式图片变成 LaTeX 序列的任务。它和普通 OCR 不是一回事:OCR 输出定长文本,公式输出是带层级结构的变长序列。标题里的 ResNet+Transformer,就是这个项目为“图像到公式序列”选的主干和序列解码器。拿到一个 zip 源码包,最先要确认的是三件事:模型入口长什么样、数据格式是什么、训练和推理脚本能不能直接跑通。下面按这条路线拆开讲,适合要复现课程设计、做公式识别毕设,或者正在给文档识别能力补数学场景的工程师。中间会给出可抄的预处理、tokenizer、训练循环代码,以及实际跑这类项目时容易翻车的一组坑。
2. ResNet+Transformer 的公式识别模型:特征图到 LaTeX 序列
2.1 为什么是这两种结构各管一段
手写公式识别的输入是一张公式图,输出是一个符号序列,比如x^2 + 1。单单用 CNN 做分类不行,因为输出不是类别而是序列;单单用 Transformer 也不行,Transformer 不擅长直接吃任意大小的图像。常见的组合方式,是让 ResNet 先做视觉特征提取,得到一组带空间位置的特征图;再把特征图拉成序列,交给 Transformer 做自注意力建模和自回归解码。前者解决“图像里有什么”,后者解决“符号按什么顺序说出来”。
为什么选 ResNet 而不是轻量 CNN?公式图像和自然场景图像不一样,它没有复杂的背景纹理,但有很多小尺寸但语义关键的符号:点、撇、上下标、根号。ResNet 的残差结构在模型加深时梯度路径稳定,预训练权重也容易拿到,用torchvision的resnet18/resnet50初始化一个 backbone 很省事。在源码级复现时我一般直接用resnet50前四层,把第五层(stride=2 的下采样)去掉,或者把 conv5 的 stride 改成 1,目的是让最终特征图的空间分辨率保持在输入的 1/8 而不是 1/32。
Transformer 侧讲的是“序列生成”。它不能直接看图片,但能吃 ResNet 送出来的“视觉 token”。它内部的自注意力会判断一个 token 和左边、右边、上下哪些 token 相关,这对公式识别很重要,因为公式的空间关系是二维的:上标在右上,下标在右下,分数线的上下各是一个子树。如果只用一维 LSTM,这种二维位置关系很难学。Transformer 因为有显式的位置编码,可以在特征里把“第几行第几列”的信息编码进去。
2.2 ResNet 侧:不要只拿最后一层特征
这里要多说一句“粗粒度/细粒度特征”。ResNet 的深层特征通道数大、语义强,但空间分辨率低,也就是说每个点对应原图 32×32 的区域,一个 8×8 的小符号在最后一层可能只占零点几个像素,基本丢了。手写公式恰恰有很多小符号,所以项目里更可靠的做法是截取 stage3/stage4 的输出做特征融合,或者像 FPN 一样把低层高分辨率特征和高层语义特征加起来再送进 Transformer。
我在代码里通常这样处理 ResNet 侧:
import torch import torch.nn as nn from torchvision.models import resnet50 class ResNetBackbone(nn.Module): def __init__(self, out_dim=512): super().__init__() base = resnet50(pretrained=True) # 取到 resnet 的 conv4 输出,保留 1/8 分辨率 self.stem = nn.Sequential( base.conv1, base.bn1, base.relu, base.maxpool, base.layer1, base.layer2, base.layer3, ) # 输出通道为 1024,尺寸为 H/8 × W/8 self.reduce = nn.Conv2d(1024, out_dim, 1) # 降维到 Transformer 的 d_model def forward(self, x): feat = self.stem(x) # [B, 1024, H/8, W/8] feat = self.reduce(feat) # [B, out_dim, H/8, W/8] b, c, h, w = feat.shape feat = feat.flatten(2).transpose(1, 2) # [B, h*w, c] return feat这段代码的逻辑是:resnet50的 layer3 输出是原图 1/8 分辨率的特征图,通道 1024。用一个 1×1 卷积把通道压到 Transformer 的d_model(比如 512),然后 flatten 成 token 序列。这样每个 token 对应原图一个 8×8 区域,保留了上下标和点的细节。参数上值得注意的一点是pretrained=True只对 ImageNet 训练有效,公式图片是灰度图,通常要把单通道复制成三通道再输入,否则 ResNet 第一个卷积层会直接报维度错。
如果你想要更接近 FPN 的多尺度融合,可以在 layer2、layer3、layer4 各引一条分支,上采样到同一尺寸后 concat。高分项目里经常会看到这种改法,效果确实比单层特征稳定,代价是显存涨一截,backbone 前向慢了 20%-30%。复现时建议先用单层跑通,再考虑融合。
2.3 Transformer 侧:编码器、解码器和 5 个必调参数
拿到视觉 token 之后,Transformer 有两条主流布置路线。
第一条路线是 Encoder-Decoder。ResNet 送出的序列进 Transformer encoder,内部自注意力做视觉 token 之间的二维关系建模;decoder 再用目标序列做 masked self-attention,逐步生成公式 token。这条路线和机器翻译最接近,很多套代码就是把 Transformer 的 encoder 输入从词向量换成 ResNet 特征。优点是工程好改,PyTorch 官方有 transformer 模块可以拼;缺点是公式识别对序列顺序比翻译更敏感。
第二条路线是 Decoder-only,也就是把视觉 token 作为前缀输入,后面接要生成的公式 token,整体一个 transformer decoder 搞定。这个做法在近两年的项目里越来越常见。手写公式的结构很复杂,比如一个\frac会引出“分支-子树”式的生成,decoder-only 能天然自回归地展开这个树。但注意:公式序列并不完全是一维树形,上标和下标在 LaTeX 里是先后写的,所以自回归顺序本身也是标注顺序,这一点训练数据一致性很重要。
我一般把 encoder 单独建模成一层可选项,先跑通 decoder。下面是核心解码器参数:
| 参数 | 常见取值 | 说明 |
|---|---|---|
| d_model | 512 | ResNet 特征降维后的通道数,同时也是注意力维度 |
| nhead | 8 | 注意力头数,太小并行性差,太大每个头的维度会碎 |
| num_layers | 4~6 | 公式识别 4 层基本够,层数上去并不一定稳 |
| dim_feedforward | 1024 或 2048 | FFN 中间层宽度,和显存直接相关 |
| dropout | 0.1 | 公式数据量不大,dropout 太低容易过拟合 |
还有两个位置编码细节容易踩坑。第一个是视觉 token 的位置编码:ResNet 输出的 token 是二维的,如果只按拉平顺序加一维位置编码,模型就不知道同一列在干什么。常见补救是做成二维位置编码:一个 H 维的位置表和一个 W 维的位置表,两者相加作为最终位置编码。第二个是输出端的位置编码,目标序列只有一维,但要注意把<sos>和<eos>处理好,否则训练时模型会在第一个时间步就预测错误起始符,损失下不去。
位置编码这块我会写成:
class PositionalEncoding2D(nn.Module): def __init__(self, d_model, max_h=64, max_w=512): super().__init__() pe_h = torch.zeros(max_h, d_model // 2) pe_w = torch.zeros(max_w, d_model // 2) pos_h = torch.arange(max_h).unsqueeze(1).float() pos_w = torch.arange(max_w).unsqueeze(1).float() div = torch.exp(torch.arange(0, d_model // 2, 2).float() * (-torch.log(torch.tensor(10000.0)) / (d_model // 2))) pe_h[:, 0::2] = torch.sin(pos_h * div) pe_h[:, 1::2] = torch.cos(pos_h * div) pe_w[:, 0::2] = torch.sin(pos_w * div) pe_w[:, 1::2] = torch.cos(pos_w * div) self.h = nn.Parameter(pe_h.unsqueeze(0), requires_grad=False) self.w = nn.Parameter(pe_w.unsqueeze(0), requires_grad=False) def forward(self, feat_h, feat_w): return self.h[:, :feat_h] + self.w[:, :feat_w]这个类生成一个 H 方向的位置基和一个 W 方向的位置基,相加得到每个 token 的位置向量。forward 里只取需要的分辨率,这样训练时给 64,推理时给 128 也能兼容。加位置编码时记得用 LayerNorm 把特征和位置信号的尺度对齐,不然模型早期容易被位置信号带偏。
3. 数据集与预处理:把 CROHME 手写公式变成能训 Transformer 的样本
3.1 先定数据:CROHME 是公式识别绕不开的基准
手写数学公式识别有一个公开基准叫 CROHME,这是文档分析与识别领域中手写数学表达式识别的评测集,CROHME 2014/2016 的离线数据是最常被引用的版本。它提供手写公式图片、对应的 LaTeX 标注以及 stroke 笔画数据。离线识别任务只用图片即可。因为标题里写的是“手写数学公式识别”,我建议复现时第一选择就是 CROHME 的离线部分,不要一上来就造自己的数据。原因是:公式识别对标注一致性要求极高,自造数据时一个人标注的\frac写法和另一个人可能不一样,模型学起来非常痛苦。
CROHME 的图是灰度图,尺寸不固定,LaTeX 标注里包含大量结构命令。如果你拿到的 zip 里没有数据,只给了数据接口,那需要去 CROHME 官方页面下载离线数据,并保证目录形式和源码里 data 路径一致。如果 zip 里已经带了数据,也要检查图片格式和标注格式是否和常见版本一致,否则后续的 tokenizer 会错位。
3.2 图像侧:先裁边,再定高缩放,最后补宽度
手写公式图片最常见的分布是:公式写在画面中间,四周有很大白边。直接放缩会让符号占的像素太少,所以要先把空白裁掉。实现的顺序是:转灰度 → 找前景像素的包围盒 → 裁边 → 等比例缩放到固定高度 → 按最大宽度 padding → 归一化。
我一般这么写:
import cv2 import numpy as np def load_formula_image(path, target_h=64, max_w=512, pad_value=0.0): img = cv2.imread(path, cv2.IMREAD_GRAYSCALE) _, binary = cv2.threshold(img, 128, 255, cv2.THRESH_BINARY_INV) ys, xs = np.where(binary > 0) x0, x1 = xs.min(), xs.max() y0, y1 = ys.min(), ys.max() crop = img[max(0, y0-4): y1+5, max(0, x0-4): x1+5] # 四周留 4px 余量 h, w = crop.shape scale = target_h / h new_w = max(1, int(round(w * scale))) resized = cv2.resize(crop, (new_w, target_h), interpolation=cv2.INTER_LINEAR) # 右端补齐到固定宽度,便于 batch 训练 padded = np.full((target_h, max_w), pad_value, dtype=np.float32) padded[:, :new_w] = resized return padded, new_w逻辑说明:先用反阈值二值化找到前景坐标,随后裁出包含公式的矩形,四周留 4 像素防止贴边符号被切掉。接着按高度 64 等比缩放,宽度随比例变化,最后右端补到 512。返回两个值:padded是模型输入,new_w是真实宽度,推理时用来截断输出或对齐位置编码。
这里有个参数容易被忽略:pad_value用 0 还是 255 要看数据标准化方式。如果后面是要减均值除方差,pad 用 0 再标准化等于“背景是均值”;如果直接输入网络,公式常用的做法是白色背景归一化到 0,黑色笔迹是负值,pad 用 0 没问题。
3.3 标签侧:LaTeX 要切成 token,不是逐字切
公式标注是 LaTeX 字符串,例如\frac{-b \pm \sqrt{b^2-4ac}}{2a}。Transformer 输出的单位是 token,不是字符,也不是单词。\frac应该作为这一个 token,b是一个字符 token,^是一个 token。建议的切分方式是先按字符串里的反斜杠命令分组:\frac、\sqrt、\pm、\times各占一个 token,字母数字和特殊符号如^ _ { }各占一个 token。因此 tokenizer 仍然比较简单,但注意统一符号写法,比如全部用\frac而不是\dfrac,不能在训练集和验证集混着来。
构建词典的代码可以这样写:
import re from collections import Counter def tokenize_tex(tex: str) -> list[str]: # 先抓命令,再抓单个字符 tokens = re.findall(r'\\[a-zA-Z]+|[a-zA-Z]|\d+|[^\s]', tex) return tokens # 遍历训练集统计词频 vocab_counter = Counter() for _, tex in train_samples: for tok in tokenize_tex(tex): vocab_counter[tok] += 1 vocab = ['<pad>', '<sos>', '<eos>', '<unk>'] + [t for t, c in vocab_counter.most_common()] tok2idx = {t: i for i, t in enumerate(vocab)} idx2tok = {i: t for t, i in tok2idx.items()}这段代码的正则先匹配“反斜杠开头的一串字母”,再匹配单个字母、连续数字和任意单个非空白字符。注意\frac在正则匹配时是\frac整体一个 token,而不是先匹配到\再匹配 f,因为第一个分支优先。下半部分用 Counter 统计频率并构建词典。
实际项目里建议给罕见命令单独保留,不能直接扔给<unk>,因为一个根号命令出错,整棵子树都会废。标签编码需要加上起始符和结束符:
def encode_tex(tex: str, tok2idx: dict[str, int], max_len=128): tokens = tokenize_tex(tex)[: max_len - 2] return [tok2idx['<sos>']] + [tok2idx.get(t, tok2idx['<unk>']) for t in tokens] + [tok2idx['<eos>']]编码结果是模型训练时 decoder 的输入序列和目标序列。decode 时去掉<sos>和<eos>再把索引转回 LaTeX。这个阶段最需要防的是训练标签和推理输出不平衡:训练时用的是 teacher forcing,每一步喂真实标签;推理时用的是上一步预测结果。如果数据里有大量漏标括号,模型生成时就会倾向于丢掉右括号。
3.4 把预处理装进 Dataset 和 collate_fn
图像预处理和标签切分做完后,还需要把它们包成 PyTorch 能直接喂的数据集。一个容易忽略的问题:如果所有图片都 padding 到全局 512 宽,短公式会造成大量无效计算。所以 collate 里一般按一个 batch 内的最大宽度动态 padding,既省显存,又不破坏长宽比。
import torch from torch.utils.data import Dataset class FormulaDataset(Dataset): def __init__(self, samples, target_h=64, max_w=512): self.samples = samples self.target_h = target_h self.max_w = max_w def __len__(self): return len(self.samples) def __getitem__(self, idx): img_path, tex = self.samples[idx] img, real_w = load_formula_image(img_path, self.target_h, self.max_w) enc = encode_tex(tex, tok2idx) return torch.tensor(img).unsqueeze(0), real_w, torch.tensor(enc) def collate_formula(batch): images, real_widths, tokens = zip(*batch) h = images[0].size(1) max_w = max(real_widths) max_len = max(t.size(0) for t in tokens) img_batch = torch.zeros(len(images), 1, h, max_w) for i, img in enumerate(images): img_batch[i, :, :, :real_widths[i]] = img[:, :, :real_widths[i]] tok_batch = torch.full((len(tokens), max_len), tok2idx['<pad>'], dtype=torch.long) for i, t in enumerate(tokens): tok_batch[i, :t.size(0)] = t # decoder 输入去掉最后一个位置,目标去掉起始符 return img_batch, tok_batch[:, :-1], tok_batch[:, 1:]这里__getitem__返回的是(图像, 真实宽度, 编码序列),collate_fn再按真实宽度动态合成 batch。tok_batch[:, :-1]作为 decoder 的输入,tok_batch[:, 1:]作为预测目标,这样<sos>和<eos>的位置严格对齐。如果之后要做推理,别忘了额外返回每个样本的真实宽度,用来生成src_key_padding_mask。
4. 训练流程、参数设置与常见问题排查
4.1 先把源码包的结构读一遍
拿到 zip 之后,我建议先看目录结构再跑,而不是直接python train.py。这类公式识别项目通常会有这几个模块:一个 data 目录放图片和标注;一个 model 目录放 backbone、transformer、位置编码;一个train.py负责训练循环,一个inference.py负责推理;还有工具脚本做可视化。先用find . -type f或tree看一遍,确认数据集路径、预训练权重路径、输出目录是否写死。
如果代码里是硬编码的绝对路径,大概率是作者在自己电脑上跑的,你需要在配置文件里改成相对路径或环境变量。常见项目会用config.yaml或 argparse 存超参数,改起来会轻松一些。最稳的第一步是:把 batch size 调小,在一个小数据集上跑一个 epoch,把训练循环和数据处理链路走通,再上全量。
4.2 训练循环:掩码是公式识别最容易写错的地方
训练时模型把“拉平的视觉 token”作为 encoder 输入,把公式 token 序列作为 decoder 输入。损失函数通常用交叉熵,计算时忽略<pad>。关键在 mask:attention 矩阵要禁止 decoder 看到未来 token。漏掉 padding mask 的典型现象是 loss 很快很低,但生成的一堆是<pad>。
核心训练步:
import torch import torch.nn as nn def generate_square_subsequent_mask(sz: int) -> torch.Tensor: mask = torch.triu(torch.ones(sz, sz) * float('-inf'), diagonal=1) return mask def train_step(model, batch, optimizer, criterion, device): img, src_key_padding_mask, tgt_in, tgt_out = batch tgt_mask = generate_square_subsequent_mask(tgt_in.size(1)).to(device) # tgt_key_padding_mask 标记 tgt_in 里是 <pad> 的位置 tgt_key_padding_mask = (tgt_in == tok2idx['<pad>']).transpose(0, 1) logits = model(img, tgt_in, tgt_mask, src_key_padding_mask, tgt_key_padding_mask) loss = criterion(logits.reshape(-1, logits.size(-1)), tgt_out.reshape(-1)) optimizer.zero_grad() loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 5.0) optimizer.step() return loss.item()逻辑说明:generate_square_subsequent_mask生成一个上三角全负无穷的矩阵,对角线及以下为 0,这样 attention 中未来位置被遮掉;tgt_key_padding_mask把 padding 位置标记成 True;loss 计算时把 batch 和序列两维合并。这里clip_grad_norm_是公式识别训练的一个隐性关键点:公式梯度里常有指示性强的符号差异,比如一个位置预测\sqrt和预测}的梯度差别非常大,不加梯度裁剪很容易在某一步把权重击穿,然后就再也回不来了。
4.3 常见问题 1:训练 loss 不降,图像输入侧出了问题
现象是 loss 一直在 8 以上徘徊,几十个 epoch 之后也没明显下降。原因是常见的数据处理错误:把三通道预训练模型拿给单通道灰度图用,输入形状不对;或者把 0-255 的灰度图直接喂进模型,数值过大导致 ResNet 的 BN 层完全混乱。解决方式是:把灰度图复制成三通道后,按 ImageNet 的 mean/std 归一化。代码里可以在 Dataset 的__getitem__里做:img = np.repeat(img[None], 3, axis=0),再(img / 255 - mean) / std。如果嫌麻烦,先打印一个 batch 里输入张量的 min、max、mean,看看是不是真的落在合理区间。
4.4 常见问题 2:训练正常,但推理输出的是“戴着括号的乱码”
现象是训练 loss 0.3,验证集也能跑,但推理输出是{ \frac { x }{ 2 } }这种嵌套全部闭合、渲染成图片却少了符号的串。原因是 vocab 里出现了形如{和}的 token,模型学会了机械配对,但没有学会结构。常见成因是数据标注不统一:训练集里\frac{1}{2}和\frac 1 2混用,导致模型把花括号当作无意义符号。解决方式是给公式做语法归一化,把不带花括号的\frac补齐;同一套数据重跑 tokenizer,保证左右花括号一定成对出现。
4.5 常见问题 3:ExpRate 一评估就崩,但 loss 很低
现象是训练 loss 掉到 0.2,验证集准确率却不高,或者突然从 60% 掉到 20%。原因是评估代码里用了 teacher forcing:把真实标签一步步喂给 decoder 看输出,模型只要学会“跟着真实标签走”就能得到很低的 loss;真正推理是自回归的,误差会累积,一旦某一步预测错,后面全错。解决方式是评估时强制用自回归推理,每个时间步取模型 prediction 作为下一步输入,而不是取真实标签。通常用生成的 sequence 和 ground truth 完全比对来计算公式级准确率 ExpRate。
4.6 常见问题 4:推理时位置编码越界
现象是训练时统一 padding 到 512,推理时来了一张更长的公式图,报维度不匹配或者位置编码越界。原因是位置编码表在初始化时写死了max_len,而输入图片宽度没限制。解决方式是把位置编码的max_h/max_w设大,比如 96×1024,训练时只取前 64×512;更保险的做法是限制输入图片最大宽度,长公式等比缩小后 padding 到固定尺寸。如果你用的是二维位置编码,还要注意 H 方向和 W 方向都做越界保护。
5. 跑通后的验证闭环:beam search、结构校验和下一步
5.1 用 beam search 换掉贪心解码
模型跑通之后,最简单的推理是每个时间步取概率最高的 token,也就是贪心解码。手写公式识别很容易在某个中间 token 出错,导致后续全错。常见做法是把解码改成 beam search,同时保留 5 条候选路径,最后按累积 log 概率挑一个。实现时不需要重写模型,只需要维护一个候选列表,在每一步对每个候选扩展 top-k 个 token,再截取 top-beam_max 条。beam size 从 1 加到 5,ExpRate 通常能涨 5-10 个点,代价是推理时间乘上大约 beam size 倍。
5.2 结构校验:比字符串比对更稳
公式识别评估不能只看字符串完全相等。有两个公式,一个写\frac{1}{2},一个写\frac12,语义一样,字符串不同。所以在验证阶段我会把两个注意点放进代码里:第一,把预测的 LaTeX 和标注统一用同一套归一化函数处理后再比对,而不是直接字符串比较;第二,对预测结果做一次括号配对和“\frac后必须有 2 个子结构”的语法检查,把明显不闭合的结果直接过滤掉。这两步能快速判断模型问题出在符号识别还是结构生成上。
5.3 一个值得投入的下一步
如果这个项目后续要往工程交付走,我会建议做一次数据增强,而不是急着换更大的 backbone。手写公式线条粗细差异大,训练时做随机腐蚀、膨胀、轻微旋转 5 度、随机裁掉上下 2% 边缘,可以让 ResNet 侧对笔迹差异更鲁棒。这套增强加在数据读取阶段,不用改模型。每次交付前,我会先在验证集上重跑一遍 beam search,再把随机抽的 20 张图渲染成 LaTeX 后人工核对一遍。形成这个习惯之后,公式识别项目翻车的概率会低很多,希望帮到你。
本文还有配套的精品资源,点击获取