
简介一份基于ResNet与Transformer架构的手写数学公式识别Python源码属于高分课程大作业项目适合深度学习、计算机视觉方向的开发者参考。项目已通过导师指导与验收代码经过严格调试可直接运行覆盖数据加载与预处理、模型搭建、训练验证及公式识别预测的完整流程。压缩包共32个文件以Python脚本为主含19个py源文件、8个pyc编译文件与2个txt字典/结果文件另附cfg和yaml配置文件整体仅87KB便于快速浏览和部署。代码按datamodule、model等功能模块组织内含编码器、解码器、位置编码、配置加载与验证脚本并附带单条公式识别结果文本可帮助理解识别管线。目前已有646人学习下载适合想跑通手写公式识别基线、在现有结构上做模块替换或二次开发的读者使用。1. 手写公式识别项目到手先看清 ResNetTransformer 在解决哪件事手写数学公式识别是把一张手写公式图片变成 LaTeX 序列的任务。它和普通 OCR 不是一回事OCR 输出定长文本公式输出是带层级结构的变长序列。标题里的 ResNetTransformer就是这个项目为“图像到公式序列”选的主干和序列解码器。拿到一个 zip 源码包最先要确认的是三件事模型入口长什么样、数据格式是什么、训练和推理脚本能不能直接跑通。下面按这条路线拆开讲适合要复现课程设计、做公式识别毕设或者正在给文档识别能力补数学场景的工程师。中间会给出可抄的预处理、tokenizer、训练循环代码以及实际跑这类项目时容易翻车的一组坑。2. ResNetTransformer 的公式识别模型特征图到 LaTeX 序列2.1 为什么是这两种结构各管一段手写公式识别的输入是一张公式图输出是一个符号序列比如x^2 1。单单用 CNN 做分类不行因为输出不是类别而是序列单单用 Transformer 也不行Transformer 不擅长直接吃任意大小的图像。常见的组合方式是让 ResNet 先做视觉特征提取得到一组带空间位置的特征图再把特征图拉成序列交给 Transformer 做自注意力建模和自回归解码。前者解决“图像里有什么”后者解决“符号按什么顺序说出来”。为什么选 ResNet 而不是轻量 CNN公式图像和自然场景图像不一样它没有复杂的背景纹理但有很多小尺寸但语义关键的符号点、撇、上下标、根号。ResNet 的残差结构在模型加深时梯度路径稳定预训练权重也容易拿到用torchvision的resnet18/resnet50初始化一个 backbone 很省事。在源码级复现时我一般直接用resnet50前四层把第五层stride2 的下采样去掉或者把 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_dim512): super().__init__() base resnet50(pretrainedTrue) # 取到 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 区域保留了上下标和点的细节。参数上值得注意的一点是pretrainedTrue只对 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_model512ResNet 特征降维后的通道数同时也是注意力维度nhead8注意力头数太小并行性差太大每个头的维度会碎num_layers4~6公式识别 4 层基本够层数上去并不一定稳dim_feedforward1024 或 2048FFN 中间层宽度和显存直接相关dropout0.1公式数据量不大dropout 太低容易过拟合还有两个位置编码细节容易踩坑。第一个是视觉 token 的位置编码ResNet 输出的 token 是二维的如果只按拉平顺序加一维位置编码模型就不知道同一列在干什么。常见补救是做成二维位置编码一个 H 维的位置表和一个 W 维的位置表两者相加作为最终位置编码。第二个是输出端的位置编码目标序列只有一维但要注意把sos和eos处理好否则训练时模型会在第一个时间步就预测错误起始符损失下不去。位置编码这块我会写成class PositionalEncoding2D(nn.Module): def __init__(self, d_model, max_h64, max_w512): 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_gradFalse) self.w nn.Parameter(pe_w.unsqueeze(0), requires_gradFalse) 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_h64, max_w512, pad_value0.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): y15, max(0, x0-4): x15] # 四周留 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), interpolationcv2.INTER_LINEAR) # 右端补齐到固定宽度便于 batch 训练 padded np.full((target_h, max_w), pad_value, dtypenp.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应该作为这一个 tokenb是一个字符 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_len128): 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_h64, max_w512): 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], dtypetorch.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。关键在 maskattention 矩阵要禁止 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), diagonal1) 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 位置标记成 Trueloss 计算时把 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, axis0)再(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 常见问题 3ExpRate 一评估就崩但 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 加到 5ExpRate 通常能涨 5-10 个点代价是推理时间乘上大约 beam size 倍。5.2 结构校验比字符串比对更稳公式识别评估不能只看字符串完全相等。有两个公式一个写\frac{1}{2}一个写\frac12语义一样字符串不同。所以在验证阶段我会把两个注意点放进代码里第一把预测的 LaTeX 和标注统一用同一套归一化函数处理后再比对而不是直接字符串比较第二对预测结果做一次括号配对和“\frac后必须有 2 个子结构”的语法检查把明显不闭合的结果直接过滤掉。这两步能快速判断模型问题出在符号识别还是结构生成上。5.3 一个值得投入的下一步如果这个项目后续要往工程交付走我会建议做一次数据增强而不是急着换更大的 backbone。手写公式线条粗细差异大训练时做随机腐蚀、膨胀、轻微旋转 5 度、随机裁掉上下 2% 边缘可以让 ResNet 侧对笔迹差异更鲁棒。这套增强加在数据读取阶段不用改模型。每次交付前我会先在验证集上重跑一遍 beam search再把随机抽的 20 张图渲染成 LaTeX 后人工核对一遍。形成这个习惯之后公式识别项目翻车的概率会低很多希望帮到你。本文还有配套的精品资源点击获取