尧图精选

ResNet+Transformer手写公式识别:原理、PyTorch代码与调优

🕒 发布时间:2026/10/2 14:15:46 📁 来源:尧图网络
简介资源包内含基于残差网络ResNet与Transformer架构的手写数学公式识别Python源码面向具备一定深度学习基础、需要完成课程设计或复现公式识别项目的学习者。方案以ResNet提取图像深层视觉特征再由Transformer自注意力机制建模公式符号的布局与依赖关系代码曾作为个人高分大作业并通过导师验收经严格调试可正常运行。压缩包共40个文件约4.21MB以19个Python脚本、8个pyc编译文件为主另含6个zbak备份、3个说明txt、2个zip补充包和config.yaml、setup.cfg等配置覆盖模型定义、数据模块、训练与测试脚本及使用说明。目前已有152人学习工程目录结构清晰适合按模块研读既能借鉴ResNet与Transformer的融合思路也可基于脚本快速复现识别效果便于后续改进或实验对比。请注意资源仅限学习交流请勿商用。1. 手写数学公式识别难在哪ResNet 负责看Transformer 负责读手写数学公式识别这个题目近几年在毕业设计和算法竞赛里出现频率很高核心诉求就一句话给一张手写公式的照片输出对应的 LaTeX 序列。基于 ResNet 与 Transformer 的组合是目前这类项目里最稳的解法——难点从来不在认单个字符而在读二维结构一张潦草的 \frac{a}{b}人一眼能看出分数线把 a 和 b 上下分开普通 OCR 却只会把它当成三个从左到右排开的字符。ResNet 把图像压成带空间信息的特征图Transformer 解码器用自注意力机制把特征图按行文顺序读成 token 序列正好一个管看一个管读。这篇文章按架构、代码、调参、排错的顺序把这个高分项目拆开讲代码基于 PyTorch思路可以直接移植到自己的数据集上。2. 编码器与解码器的分工为什么 ResNet Transformer 能处理二维结构的公式2.1 ResNet 编码器粗粒度特征管布局细粒度特征管符号ResNet 在这套方案里不是分类器而是纯粹的视觉特征提取器。常见做法是加载 ImageNet 预训练的 resnet18 或 resnet34砍掉最后的全局池化和全连接层只保留卷积部分。为什么不用更深的 resnet50公式图像本身纹理简单深度网络在小数据集上容易过拟合推理也慢还会把特征图分辨率压低。特征图的粒度是第一个关键取舍。resnet18 的 stage4 输出步长为 32stage3 输出步长为 16。步长 32 的特征图感受野大、序列短属于粗粒度特征看全局布局够用但 \dot 上面的小点、下标里的小数字会糊成一团步长 8 能留住细粒度细节可序列长度翻四倍自注意力的显存开销是平方级增长。下面这张表是我在不同步长下对 160×640 输入的实测感受输出步长特征图尺寸解码器序列长度细粒度符号表现显存压力325×20100差上标和下标边界容易粘连低1610×40400中等偏上能分辨 \frac 上下层中820×801600好小符号清晰高batch 稍大就 OOM我一般固定用 stride 16 的输出。序列长度几百个 token自注意力还能跑得动细粒度空间信息——上标在右上、下标在右下、分数线上下分层——也保留得住。如果项目里想上 FPN 把 stage3 和 stage2 的细粒度特征融合进来位置编码也得跟着设计因为不同层的分辨率不一样那是另一个复杂度不建议第一次做就加。2.2 Transformer 解码器掩码自注意力保证按顺序生成交叉注意力对齐图像区域Transformer 解码器做的事情可以理解为边看图像边写 LaTeX。每一层包含两个注意力子层掩码自注意力masked self-attention负责维护已经写过的内容之间的关系交叉注意力cross-attention负责从 ResNet 特征图里按需取信息。生成是自回归的上一步输出的 token 拼进序列再预测下一步。掩码自注意力是整个解码器最容易写错的地方。训练时如果允许当前位置看到后面的 token模型就能作弊——直接复制下一个正确答案。这样 loss 会掉得很快但推理时没有未来信息beam search 的输出就会变成一堆重复括号。标准做法是准备一个上三角全 True 的布尔矩阵保证第 t 个位置只能看到 0 到 t 的内容。这个矩阵在第三章代码里会具体给出来。交叉注意力是这套方案里最值钱的机制每个输出 token 都对应一张二维注意力热图标明生成这个符号时模型在看图像的哪个区域。比如生成 \frac 的分子 token 时热图应该落在图像上半部分生成分母 token 时应该落在分数线下方。这是后续调试、可视化和判断模型有没有学到结构的最直接依据比盯 loss 曲线靠谱得多。2.3 为什么不直接用 CRNNCTC三种方案的对比有更轻的方案但它们各自有硬伤。CNNCTC 把公式当成一维文本流识别工程简单、速度快可分数、根号、矩阵这种二维结构天然表达不了结构错乱率很高CNNRNN 编码解码能输出带结构的 LaTeX但 LSTM 在长序列上信息衰减明显公式动辄几十个 token训练串行也很慢ViT 这类把图像直接切 patch 进 Transformer 的思路表达能力好但对数据量的要求比 ResNet 高一个量级手写公式这种小样本任务很容易欠拟合。ResNet 提供廉价的局部视觉先验——边缘、拐角、笔顺纹理——Transformer 负责全局结构建模两者互补是这类项目里收益最高的组合。高分项目的评分点往往也在结构正确率上这一条选型理由就值得写进报告。3. 从 tokenizer 到 beam search一套可运行的 PyTorch 代码骨架3.1 数据准备图像、LaTeX 标签和 tokenizer公开的 CROHME 数据集是手写公式识别的事实标准包含手写公式图像和对应 LaTeX 标签没有数据的话常见做法是用 LaTeX 渲染合成公式图像做预训练再用少量手写样本微调。数据准备里最容易被低估的一步是 tokenizer——LaTeX 标签里同类写法太多\dfrac 和 \frac 等价、\left( 和 ( 等价如果不先归一化模型要把一套语义学两遍词表还白白膨胀。import re def normalize_formula(f): LaTeX 标签归一化消除等价写法减少词表膨胀 f f.replace(r\dfrac, r\frac) f f.replace(r\displaystyle, ) f re.sub(r\s, , f) # 去掉空白避免同一公式两种写法 f re.sub(r\\left|\\right, , f) # 去掉可省的左右定界符 return f class FormulaTokenizer: 把 LaTeX 公式切成 token 序列并维护 pad/sos/eos 特殊 token TOKEN_RE re.compile(r\\[a-zA-Z]|[a-zA-Z0-9]|[{}_^\-()/.,;!]|.) def __init__(self, formulas, max_vocab200): self.tokens [pad, sos, eos, unk] freq {} for f in formulas: for tok in self.TOKEN_RE.findall(normalize_formula(f)): freq[tok] freq.get(tok, 0) 1 for tok, _ in sorted(freq.items(), keylambda x: -x[1]): if len(self.tokens) max_vocab: break self.tokens.append(tok) self.stoi {t: i for i, t in enumerate(self.tokens)} self.itos {i: t for i, t in enumerate(self.tokens)} self.pad self.stoi[pad] self.sos self.stoi[sos] self.eos self.stoi[eos] def encode(self, f, max_len128): ids [self.sos] for tok in self.TOKEN_RE.findall(normalize_formula(f)): ids.append(self.stoi.get(tok, self.stoi[unk])) if len(ids) max_len: break ids.append(self.eos) return ids这个 tokenizer 的逻辑说明正则表达式把 \frac、\sum 这类 LaTeX 命令整体切成一个 token把单个数字和字母各切一个 token花括号、上下标符号单独成 tokenmax_vocab200 是词表上限训练集里出现次数最少的生僻符号会落到 对公式识别影响不大。encode 的返回值是sos...eos的 id 列表训练时会把最后一个 eos 拆出去当预测目标。参数说明max_vocab 太小会把 \sqrt、\operatorname 这类命令挤成 导致结构错误太多又让输出层参数暴增。手写公式场景 200 到 300 之间是常见区间。数字逐位切分是刻意的设计——模型永远能拼出训练集里没出现过的数字而不是死记硬背。3.2 模型定义ResNet 提特征Transformer 解码器出序列模型主体分四块ResNet 卷积部分提特征、1×1 卷积压通道、Transformer 解码器生成序列、线性层映射到词表。这里有一个新手容易踩的点ResNet 输出是二维特征图Transformer 只吃序列所以要把特征图按行展开成序列展开顺序必须和位置编码的构造顺序一一对应否则上下结构会学反。import math import torch import torch.nn as nn from torchvision import models class FormulaModel(nn.Module): def __init__(self, vocab_size, d_model256, nhead8, num_layers4, dropout0.1, backboneresnet18): super().__init__() # ResNet 去掉池化和全连接只留卷积部分 resnet getattr(models, backbone)(weightsmodels.ResNet18_Weights.IMAGENET1K_V1) self.cnn nn.Sequential(*list(resnet.children())[:-2]) # 输出 512 通道 self.proj nn.Conv2d(512, d_model, kernel_size1) # 压到 d_model self.embed nn.Embedding(vocab_size, d_model) # 解码器词嵌入 decoder_layer nn.TransformerDecoderLayer( d_modeld_model, nheadnhead, dim_feedforward1024, dropoutdropout, batch_firstTrue) self.decoder nn.TransformerDecoder(decoder_layer, num_layers) self.out nn.Linear(d_model, vocab_size) self.d_model d_model self.pos2d None # 二维位置编码按特征图尺寸惰性生成 def forward(self, img, tgt, tgt_mask, tgt_pad_maskNone): feat self.cnn(img) # (B, 512, H, W) feat self.proj(feat) # (B, d_model, H, W) B, C, H, W feat.shape mem feat.flatten(2).transpose(1, 2) # (B, H*W, d_model)行优先展开 if self.pos2d is None or self.pos2d.shape[1] ! H * W: self.pos2d self._make_pos(H, W, C).to(img.device) mem mem self.pos2d # 给每个空间位置加二维编码 tgt self.embed(tgt) * math.sqrt(self.d_model) tgt tgt self._pos_1d(tgt.size(1), C).to(tgt.device) out self.decoder(tgt, mem, tgt_masktgt_mask, tgt_key_padding_masktgt_pad_mask) return self.out(out) # (B, T, vocab_size) staticmethod def _pos_1d(T, d_model): 一维正弦位置编码用于解码器序列位置 div torch.exp(torch.arange(0, d_model, 2) * (-math.log(10000.0) / d_model)) pos torch.arange(T).unsqueeze(1) pe torch.zeros(T, d_model) pe[:, 0::2] torch.sin(pos * div) pe[:, 1::2] torch.cos(pos * div) return pe.unsqueeze(0) # (1, T, d_model) staticmethod def _make_pos(H, W, d_model): 二维正弦位置编码行坐标和列坐标各占一半维度 half d_model // 2 div torch.exp(torch.arange(0, half, 2) * (-math.log(10000.0) / half)) row torch.arange(H).unsqueeze(1) * div col torch.arange(W).unsqueeze(1) * div pe_row torch.zeros(H, half) pe_row[:, 0::2] torch.sin(row) pe_row[:, 1::2] torch.cos(row) pe_col torch.zeros(W, half) pe_col[:, 0::2] torch.sin(col) pe_col[:, 1::2] torch.cos(col) # 行编码每行重复 W 次列编码按行展开拼成 d_model 维 pe torch.cat([ pe_row.repeat_interleave(W, dim0), pe_col.unsqueeze(0).expand(H, W, -1).reshape(-1, half) ], dim1) return pe.unsqueeze(0) # (1, H*W, d_model)参数说明d_model256 是序列特征的宽度nhead8 必须能整除 d_modeldecoder 层数 4 到 6 之间。dim_feedforward1024 是前馈网络宽度越大表达能力越强也越容易过拟合。位置编码的构造逻辑是核心——位置编码和特征图是相加关系比例不需要额外调嵌入层乘 sqrt(d_model) 是为了让 embedding 和位置编码的量级匹配这是 Transformer 原论文的写法。3.3 训练循环teacher forcing 与 mask 的正确姿势训练用 teacher forcing每一步都用标准答案的前缀作为解码器输入预测下一个 token。这样收敛快但前提是两个 mask 必须写对——因果 mask 保证只看过去padding mask 保证 pad 位置不参与注意力、也不参与损失计算。def collate_batch(batch, tokenizer, max_len128): imgs, tgt_in, tgt_out zip(*batch) imgs torch.stack(imgs) max_t max(len(t) for t in tgt_in) tgt_in_ids torch.full((len(batch), max_t), tokenizer.pad, dtypetorch.long) tgt_out_ids torch.full_like(tgt_in_ids, -100) # -100 会被损失函数忽略 for i, (ti, to) in enumerate(zip(tgt_in, tgt_out)): tgt_in_ids[i, :len(ti)] torch.tensor(ti[:max_len]) tgt_out_ids[i, :len(to)] torch.tensor(to[:max_len]) return imgs, tgt_in_ids, tgt_out_ids def train_step(model, batch, tokenizer, opt, criterion, device): img, tgt_in, tgt_out batch img, tgt_in, tgt_out img.to(device), tgt_in.to(device), tgt_out.to(device) T tgt_in.size(1) # 因果 maskTrue 表示该位置不允许被看到上三角遮住未来 causal torch.triu(torch.ones(T, T, dtypetorch.bool, devicedevice), diagonal1) pad_mask (tgt_in tokenizer.pad) # pad 不参与注意力 logits model(img, tgt_in, causal, tgt_pad_maskpad_mask) V logits.size(-1) loss criterion(logits.view(-1, V), tgt_out.view(-1)) opt.zero_grad() loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 2.0) # 梯度裁剪防训练震荡 opt.step() return loss.item()逻辑说明causal 矩阵用 torch.triu(..., diagonal1) 生成第 i 行的 True 全部落在 i 之后的列上decoder 在第 i 步就看不到 i1 及之后的 token。tgt_out 里 pad 位置填 -100配合 CrossEntropyLoss 的 ignore_index-100 参数模型预测 pad 不产生损失。参数说明梯度裁剪阈值 2.0 是 Transformer 常见设置不裁剪的话前几步 loss 容易跳动。优化器我一般用 AdamW初始学习率 3e-4 左右配合 warmup 1000 到 2000 步——Transformer 在训练前期梯度方向不稳定先用小学习率走稳再逐步放大。3.4 推理beam search 让候选序列按概率排序训练完不能直接贪心解码。贪心只取每步概率最高的 token一步选错后面全错公式结构是链式的错一个括号整句报废。beam search 同时保留 top-k 个候选序列每步扩展后按累计概率排序截断最后取总分最高的。torch.no_grad() def beam_search(model, img, tokenizer, beam5, max_len128, devicecuda): beam search 解码每步保留概率最高的 beam 个候选序列 model.eval() img img.to(device) sos, eos tokenizer.sos, tokenizer.eos beams [(0.0, [sos])] # (累计log概率, token序列) for _ in range(max_len): cands [] for score, seq in beams: tgt torch.tensor([seq], devicedevice) T tgt.size(1) causal torch.triu(torch.ones(T, T, dtypetorch.bool, devicedevice), diagonal1) logits model(img, tgt, causal)[0, -1].log_softmax(-1) top logits.topk(beam) for v, lp in zip(top.indices.tolist(), top.values.tolist()): cands.append((score lp, seq [v])) beams sorted(cands, keylambda x: x[0], reverseTrue)[:beam] if all(s[-1] eos for _, s in beams): break best max(beams, keylambda x: x[0] / len(x[1])) # 长度归一化防倾向短句 return [tokenizer.itos[i] for i in best[1][1:]]逻辑说明每条候选序列维护累计 log 概率和 token 列表每步对每个候选做一次前向扩展 beam 个新 token全部合并后按概率排序只留前 beam 个。长度归一化是因为 log 概率是负数累加序列越长分越低直接取 max 会倾向短句子除以长度后更公平。参数说明beam5 是性价比最高的档位beam 加到 10 提升有限推理时间翻倍。这个实现按单样本写方便理解批量推理时把 beam 个候选打包成一个 batch 前向速度能快一个量级但代码复杂度也上一个台阶。4. 超参数与评估指标把跑通调到高分4.1 一组能直接起步的超参数下面的参数不是拍脑袋是这类项目反复试下来比较稳的一组起点。特征图尺寸由输入图像决定d_model 和 head 数绑在一起学习率要配合 warmup 调整batch 大小受显存限制。先跑通再逐项调不要一上来全改。参数建议值说明图像高度 / 宽度160 / 640固定高度宽度按比例缩放后 pad 到 640d_model256128 容易欠拟合512 需要更大数据量nhead8必须整除 d_modeldecoder 层数4 ~ 64 层先跑通6 层提精度学习率3e-4AdamW配 warmupwarmup 步数1000 ~ 2000Transformer 前期需要热启动batch size16 ~ 32以不 OOM 为先beam 宽度5推理时用5 以上收益递减label smoothing0.1缓解过拟合对结构 token 尤其有效图像尺寸这块多说一句直接 resize 到 640×160 会把宽高比拉变形分数线的弧度会被压平我一般固定高度 160、宽度按比例缩放不够 640 就在右边 pad 成白色训练时随机裁切一点偏移让模型见过分数线不在正中间的情况。4.2 评估指标算识别对了有三种口径只看 loss 曲线看不出模型好不好用。公式识别通常报三个指标token 准确率、公式级准确率和 EXP Rate。token 准确率预测 token 和标准 token 逐位相等的比例最容易刷但结构错一个 token 整句就废了所以高分项目的报告里重点看后两个。def compute_metrics(pred_tokens, gt_tokens): 返回 token 准确率和公式级准确率 correct sum(1 for p, g in zip(pred_tokens, gt_tokens) if p g) token_acc correct / max(len(gt_tokens), 1) expr_acc 1.0 if pred_tokens gt_tokens else 0.0 # 严格相等太苛刻等价写法统一交给 normalize_formula 处理 return token_acc, expr_acc逻辑说明公式级准确率是严格全等pred 和 gt 都先经过 normalize_formula 归一化把 \dfrac、多余空格这类等价差异抹平否则模型明明识对也会被记成错误。EXP Rate 是结构等价准确率需要把 token 序列解析成表达式树再比较实现起来要引入公式解析库第一次做项目可以先用严格相等兜底报告里说明口径即可。参数说明这个函数在验证集上逐条调用最后取平均。公式级准确率在 CROHME 类数据集上基线模型 50% 到 60% 是常态beam search 和位置编码调对之后能到 70% 以上再往上拼的是数据和模型容量。4.3 位置编码怎么计算一维不够公式是二维结构Transformer 原论文里的位置编码是给一维文本设计的PE(pos, 2i)sin(pos/10000^(2i/d))PE(pos, 2i1)cos(pos/10000^(2i/d))把位置序号编码成不同频率的正弦波。解码器生成 LaTeX 是一维序列用这个没问题第三章代码里的 _pos_1d 就是标准实现。但 ResNet 特征图是二维的。flatten 成序列之后第 r 行第 c 列的位置丢了在第几行的信息——如果只给一个一维位置序号第 2 行第 1 列和第 1 行第 W 列在序列里是相邻的模型很难区分上下结构\frac{1}{2} 被识别成 \frac{2}{1} 就是这么来的。解决办法是二维位置编码行坐标和列坐标各用一半维度做正弦编码拼成 d_model 维向量再加到特征图上。第三章代码里的 _make_pos 就是完整实现这里的计算逻辑是行编码每一行的位置 r 编码成 half 维正弦向量平铺到该行的每个列位置列编码每一列的位置 c 也编码成 half 维正弦向量按行展开拼到每个位置组合后每个空间位置拿到一个 d_model 维向量既知道自己在第几行也知道在第几列。这个设计的价值在于公式特有的空间关系——分数的分子在上面、分母在下面上标在右上角、下标在右下角——全部靠行方向的位置编码表达。位置编码不是学出来的是现算的所以输入图像尺寸变了也能自适应不需要重新训练。5. 避坑手册手写公式识别最常翻车的 5 个排查现场5.1 现象训练 loss 降得很顺beam search 输出一堆重复括号原因decoder 的因果 mask 写错了。最常见是 torch.triu(..., diagonal0) 把对角线也保留当前位置能看到自己或者 padding mask 没加模型学到看到 pad 就跟着输出 pad。teacher forcing 阶段有标准答案兜底错 mask 不容易暴露一到推理就全现形。这种问题属于典型的代码看起来对、跑起来错。解决按第三章代码核对 causal 矩阵确认 diagonal1训练前用固定样本打印 mask 形状人工检查第 i 行哪些位置是 True。我习惯把 mask 可视化输出一次再开训练省掉一整天的无效迭代。5.2 现象批量训练 OOMbatch 减到 4 还是爆显存原因显存大头不是 CNN是注意力矩阵。输入 160×640 的图像经 stride 16 后是 10×40 的特征图展平成 400 个位置解码器序列长度 128 时每个注意力头的矩阵是 128×400再乘 8 个头、乘 6 层、乘 batch 和 beam 宽度这是宽扁矩阵的爆炸式增长。解决先把图像宽度从 640 降到 480序列长度直接少四分之一decoder 层数从 6 降到 4推理时 beam 分批跑不要一次性展开全部候选。显存还是不够就换用优化后的注意力实现但先确认训练环境能装得上再改代码。5.3 现象符号全对但 \frac 的分子分母位置对调原因绝大多数情况是二维位置编码没生效。常见是 flatten 展开顺序和位置编码的行列顺序不一致——特征图按行展开位置编码却按列展开上下信息完全错乱或者只套了一维位置编码行方向的信息根本没进模型。解决打印 mem 里几个固定位置的编码向量确认行分量随行号单调变化、列分量随列号单调变化。更直接的验证是解码时把交叉注意力热图画出来看分子 token 是否注视图像上半区域。这类问题用注意力可视化定位最快不用猜。5.4 现象验证集准确率上不去加数据也没用原因tokenizer 和标签没归一化。同一个公式在数据集里可能同时存在 \dfrac 和 \frac、多余花括号、\left( 和 ( 这些等价写法模型要学两套映射容量被稀释。另一个隐藏雷区是某个生僻命令全部落到 比如 \operatorname、\boldsymbol模型每次遇到都乱猜。解决数据加载阶段统一调用 normalize_formula评估时对预测和标准答案做同样处理。检查词表里 的出现频率如果某个结构命令频繁落 说明 max_vocab 设小了或者正则表达式没覆盖到这类命令。5.5 现象同一个模型换随机种子公式级准确率波动三五个点原因训练脚本没固定随机种子。ResNet 预训练权重、dropout、数据 shuffle、优化器的初始化每个环节都有随机性Transformer 对初始状态敏感微小扰动在结构预测上会被放大。这种问题最气人因为代码一行没改。解决训练脚本开头固定 torch.manual_seed、numpy.random.seed、random.seed打开 cudnn.deterministiccheckpoint 按验证集表现保存而不是最后一个 epoch。这样换机器复现时至少排除随机性这个变量。6. 进阶用交叉注意力验证模型到底在看哪训练跑通后第一件事不是急着调参数而是验证注意力。把最后一层解码器的交叉注意力权重捞出来reshape 回特征图尺寸叠在原图上就能看到生成每个 token 时模型在看图像的哪个区域。比如生成上标 ^ 时热图应该集中在公式右上角生成分母 token 时热图应该落在分数线下方。注意力热图是判断模型有没有学到结构的直接证据比盯着 loss 曲线猜有用得多。attn_map [] def hook(module, args, out): # 兼容不同 PyTorch 版本out 可能是 (out, weights) 或只有 out if isinstance(out, tuple): attn_map.append(out[1].detach()) # (B, T, H*W) 或 (B*heads, T, H*W) model.decoder.layers[-1].multihead_attn.register_forward_hook(hook)捞出来后对多个 head 取平均reshape 成 (T, 特征图高, 特征图宽)找目标 token 对应的一行画热图。heads 之间分歧大本身就说明模型不确定这也是调整后续训练重点的信号。第二个技巧是数据增强。手写样本不足时常见做法是用 LaTeX 渲染合成数据预训练再用弹性形变增强真实手写样本模拟笔迹抖动。参数上 alpha 取 4 到 8、sigma 取 0.5 到 1 比较稳——alpha 控制形变强度sigma 控制平滑范围。增强要克制alpha 超过 10 会把数字 1 和 7 扭曲成同一形状效果适得其反把增强前后的样本并排看一遍就能发现退化。我自己现在跑这类项目的固定流程是小数据集上先跑几十个 epoch出注意力热图确认模型盯对了位置、上下结构没有看反再放开全量数据和 beam search 调参。这个习惯帮我省掉了无数个训练半天才发现是数据或者 mask 错了的晚上。希望帮到你。本文还有配套的精品资源点击获取
上一篇/下一篇内容由系统自动关联 返回资讯列表 →