简介:本资源是一套基于ResNet与Transformer混合架构的手写数学公式识别Python实现,面向深度学习初学者与计算机视觉方向进阶学习者,解决手写公式图像到LaTeX序列的端到端识别问题,适用于教育数字化、智能阅卷、学术笔记OCR等场景。压缩包共40个文件,含19个核心Python源码(覆盖数据加载、ResNet特征提取、Transformer编码器/解码器、位置编码、训练验证全流程)、6个备份文件(.zbak)、3个说明类文本及配置文件(config.yaml、setup.cfg等),整体大小为4.21MB,目录结构模块化清晰,含datamodule、model、utils等规范子包。已有151人学习下载,代码经严格调试可直接运行,包含完整训练脚本(train.py)、测试脚本(test_bttr.py)、词表构建(vocab.py)及预处理工具,附带详细说明文档与单样本识别结果示例,便于理解多模态特征融合设计与序列生成逻辑。 手写数学公式识别这个题目,算是我见过的课程设计里少有的既能卷技术深度、又具备实际应用价值的选题。你想想,OCR领域常规的手写数字识别、印刷体文字识别,网上教程一抓一大把,照着跑通一个 LeNet 或者 CRNN 就算交差了。但手写数学公式不一样,它天生自带难度——符号之间不仅有左右顺序,还有上下标、分式结构、根号嵌套、矩阵布局,二维空间关系极其复杂。正因如此,把这个题目做出来并且做出效果,在课程答辩、竞赛评审里都是非常亮眼的加分项,这也是为什么这类项目经常被冠以“高分项目”的原因。
这篇博客,我直接把整个项目的思路、代码结构、训练细节、踩坑实录全部摊开来讲。项目本身采用 ResNet 作为视觉编码器、Transformer 作为序列解码器,用端到端的方式完成“公式图像 → LaTeX 序列”的映射。无论你是准备做毕业设计、课程项目,还是纯粹想系统掌握 CNN + Transformer 在视觉任务中的应用,这篇文章都能给你一份可以直接复现的参考路径。我会尽量把每个选择的“为什么”也讲清楚,而不是只丢一段能跑的代码。
1. 整体设计与技术选型思路
1.1 为什么公式识别不能套用普通 OCR 方案
先聊聊项目设计的第一步:摸清楚问题的本质。手写数学公式识别和普通文本行识别,最核心的区别在于“结构歧义”和“二维布局”。比如一个分式 \frac{a}{b},在图像上并不是 a、/、b 这种线性排列,而是 a 在上、b 在下、中间横线贯穿。如果强行用 CRNN 这种“CNN 提特征 + RNN 建模序列”的架构去识别,模型很难学到“从上往下读”这种结构关系,效果会很惨。
另外,数学公式里很多符号在视觉上非常相似。手写的“1”和“/”、“0”和“o”、“x”和“×”,人眼都常常需要结合上下文判断,更别说机器了。这种情况下,模型必须具有极强的上下文建模能力,能够根据前后符号、甚至整个公式的语义来消除歧义。RNN 类模型虽然也能建模序列,但长距离依赖捕捉能力有限,训练效率也低。这给了 Transformer 上场的机会。
1.2 为什么选 ResNet + Transformer 这套组合
项目标题里直接点名的两个模型,不是随便凑在一起的。ResNet 负责“看”,Transformer 负责“想”。
ResNet 在 2015 年提出之后,几乎成了视觉特征提取的默认底座。它通过残差连接解决了深层网络退化问题,让我们可以把网络堆到 50 层、101 层甚至更深,提取出足够高层、足够抽象的视觉特征。对比 VGG 那种纯堆卷积的结构,ResNet 不仅参数效率更高,梯度传播也更顺畅。在手写公式这种细节丰富、噪声较多的图像上,ResNet 强大的特征表达能力非常关键。
Transformer 则彻底改变了序列建模的格局。它抛弃了 RNN 的递归结构,完全基于自注意力机制,能够直接建模序列中任意两个位置之间的关系。解码器部分还可以通过 Masked Self-Attention 保证生成时的因果性,配合 Cross-Attention 实现对编码器输出特征的动态关注。这种“全局感知”能力,用在公式结构重建上再合适不过。
这个编码器-解码器框架,本质上就是“CNN 负责把图像变成语义特征序列,Transformer 负责把这些特征一步步解码成目标序列”。整套架构在图像描述(Image Captioning)、手写识别、数学公式识别等任务上都验证过,成熟度和可复现性都很高。
1.3 方案选型的心得体会
做这个项目时,我也对比过其他方案。比如两阶段方案:先做符号检测,再用图匹配或者规则引擎去分析结构。这种方案的问题在于,每个环节的误差会累积,且规则引擎面对手写变体时非常脆弱。还有直接上大模型微调,比如用多模态大模型做 few-shot 识别,效果可能不错,但对硬件要求高,也不利于课程答辩时讲清楚原理。
ResNet + Transformer 这套组合,恰好卡在一个很舒服的位置:理论基础扎实、代码实现不复杂、训练效率高、效果有保障。而且导师或评审老师一看这个架构,就知道你确实理解了现代深度学习的主流范式,提问环节也容易应对。如果你还能讲清楚 Positional Encoding、Beam Search、Teacher Forcing 这些细节,高分几乎是板上钉钉的事。
2. 数据准备与预处理
2.1 数据集选型与获取
做任何深度学习项目,数据都是第一个拦路虎。手写数学公式识别领域最常用的公开数据集是 CROHME(Competition on Recognition of Online Handwritten Mathematical Expressions)。它有离线版和在线版,我们这里用的是离线图像版。
CROHME 数据集的特点是:公式由众多书写者手写采集,风格差异大,包含符号种类约 100 多个,公式结构覆盖分式、根号、上下标、求和符号等常见类型。训练集大概有 8000 多个公式样本,测试集约 1000 个,数据量不大,但对训练一个端到端模型来说基本够用。如果实验室条件允许,还可以自己扩展一些样本,比如用触控板或者数位屏采集自己手写的公式,增强模型的泛化能力。
下载数据集时要注意文件结构。CROHME 官方提供的是 InkML 格式的笔迹数据,而我们需要的是渲染好的图像。有的工具包会提供离线渲染脚本,但如果你不想折腾,直接找社区处理好的图像版本会更高效。Kaggle 上有人整理过 CROHME 的图片形式数据集,格式为“图像 + LaTeX 标签”的配对文件,用起来非常顺手。
2.2 图像预处理细节
拿到图像后,不能直接丢给网络,需要经过几步预处理。第一步是灰度化,公式图像本身没有颜色信息,灰度图足够表达。第二步是缩放,考虑到公式图像的长宽比例差异很大,不能简单粗暴地 resize 成正方形,否则会导致符号严重变形。
我的做法是:先将图像按比例缩放,使长边不超过 256 像素,然后使用 padding 将图像补成 256×256 的正方形,padding 区域用白色填充(因为公式通常是黑字白底)。这样既保留了符号的纵横比,又满足了网络输入尺寸固定的要求。
对应到 PyTorch 代码,大概是这样的思路:
import cv2 import torchvision.transforms as T def preprocess_image(img_path, target_size=256): # 读取灰度图 img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) h, w = img.shape # 长边缩放到 target_size scale = target_size / max(h, w) new_h, new_w = int(h * scale), int(w * scale) img = cv2.resize(img, (new_w, new_h), interpolation=cv2.INTER_AREA) # 白色底 padding 到正方形 canvas = np.ones((target_size, target_size), dtype=np.uint8) * 255 y_offset = (target_size - new_h) // 2 x_offset = (target_size - new_w) // 2 canvas[y_offset:y_offset+new_h, x_offset:x_offset+new_w] = img # 转为 Tensor 并归一化 tensor = T.ToTensor()(canvas) # 值域 [0, 1] tensor = T.Normalize(mean=[0.5], std=[0.5])(tensor) # 转换到 [-1, 1] return tensor这里有一点值得提:归一化参数 mean=0.5, std=0.5 意味着把像素从 [0, 1] 映射到 [-1, 1],这是很多预训练视觉模型用到的标准方式。如果你打算用 ImageNet 上预训练的 ResNet,也可以沿用 ImageNet 的 mean/std,但那样需要把图像转成三通道。我实际测试下来,灰度图单通道丢给 ResNet 的前几层(需要自行修改输入通道),效果并不差,而且省显存。
2.3 标签编码与词表构建
公式的标签是 LaTeX 字符串,例如\frac { a } { b } + c ^ { 2 }。我们无法直接用字符串做损失计算,需要先构建一个词表(Vocabulary),把每个“符号”或“token”映射成整数 id。
注意这里的分词粒度。我的经验是按“原子符号”切分,而不是按空格切分。因为 LaTeX 中有些命令是整体语义,比如\frac是一个完整的令牌,中间不能拆开;而花括号{}只是结构标记,可以单独作为 token。切分词表后,需要给序列两端加上特殊标记:<sos>(序列开始)和<eos>(序列结束)。同时也要定义<pad>标记,用于 batch 内序列对齐。
词表构建的逻辑如下:
def build_vocab(sequences, min_freq=1): freq = {} for seq in sequences: tokens = tokenize_latex(seq) for t in tokens: freq[t] = freq.get(t, 0) + 1 vocab = ['<pad>', '<sos>', '<eos>', '<unk>'] for token, count in sorted(freq.items(), key=lambda x: -x[1]): if count >= min_freq: vocab.append(token) token2id = {t: i for i, t in enumerate(vocab)} id2token = {i: t for t, i in token2id.items()} return token2id, id2token词表规模一般控制在 200 个 token 以内,不需要很大。如果某些生僻符号在训练集里出现次数太少,直接映射到<unk>就行,硬塞进词表只会让模型过拟合到噪声上。
2.4 数据增强(容易忽视但很重要)
手写数据集样本量不大,很容易过拟合。我建议做两类增强:仿射变换和噪声扰动。
仿射变换包括小幅度的旋转(±5°)、缩放(0.95~1.05)、平移(±5%)。注意旋转角度不能太大,否则公式的语义结构会被破坏,比如“分式横线”转成斜线就麻烦了。噪声扰动可以用高斯噪声、笔画腐蚀/膨胀等手段,模拟不同笔迹的墨迹差异。
PyTorch 里可以用torchvision.transforms.RandomAffine实现。如果你用的是 Albumentations 库,它支持对图像做更复杂的增强,操作也很简便:
import albumentations as A train_transform = A.Compose([ A.RandomAffine(rotate=(-5, 5), translate_percent=(-0.05, 0.05), scale=(0.95, 1.05), p=0.5), A.GaussNoise(var_limit=(10.0, 30.0), p=0.2), A.RandomBrightnessContrast(brightness_limit=0.05, contrast_limit=0.05, p=0.2), ])这里我踩过一个坑:一开始我把增强加在验证集上,结果验证指标一直上不去,后来才发现是验证集也被随机旋转了。记住,增强只应该作用于训练集,验证集和测试集保持原始图像即可。
3. 模型结构核心实现
3.1 整体网络框架
模型的整体结构是标准的编码器-解码器架构:
- 编码器:ResNet(可以加载预训练权重),输入是 3×256×256 或 1×256×256 的图像,输出是一组特征图,尺寸为 C×H×W(比如 512×8×8)。
- 特征序列化:将特征图展平为 H×W 个位置,每个位置对应一个长度为 C 的特征向量,形成“视觉 token 序列”。
- 解码器:Transformer Decoder,输入是目标序列(训练时)或已生成的 token(推理时),通过 Self-Attention 和 Cross-Attention 逐步生成下一个 token。
这个设计的巧妙之处在于,特征图中的每个空间位置都可以被理解为“图像中的一个局部区域”,Transformer 能通过注意力机制自动决定当前应该关注图像的哪个区域。比如生成\frac之后,解码器会更关注分子和分母所在的位置。这比固定规则可靠得多。
3.2 ResNet 编码器的实现与改造
ResNet 部分可以直接用torchvision.models里现成的模型。但要注意几点修改:
- 输入通道:原版 ResNet 是 3 通道输入,如果我们的图像是单通道灰度图,需要把
conv1改成nn.Conv2d(1, 64, ...)。 - 去掉最后的全连接层和平均池化:我们只需要特征图,不需要分类结果。
- 调整输出步长:原版 ResNet 的最终特征图是输入尺寸的 1/32。对于 256×256 的输入,得到的特征图是 8×8,空间分辨率偏低。我建议把最后一个 stage 的 stride 从 2 改为 1,并用空洞卷积保持感受野,这样特征图可以到 16×16,细节信息更丰富。
实现示例:
import torch.nn as nn from torchvision import models class ResNetEncoder(nn.Module): def __init__(self, in_channels=1): super().__init__() resnet = models.resnet18(pretrained=True) # 修改第一个卷积层适配单通道 self.conv1 = nn.Conv2d(in_channels, 64, kernel_size=7, stride=2, padding=3, bias=False) 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 # 这个 1x1 卷积把 ResNet 的输出通道映射到 Transformer 的 d_model self.proj = nn.Conv2d(resnet.layer4[-1].conv2.out_channels, d_model, kernel_size=1) def forward(self, x): x = self.conv1(x) x = self.bn1(x) x = self.relu(x) x = self.maxpool(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 return x修改最后一个 stage 的 stride 需要额外小心。最简单的做法是用resnet.layer4中每个 Bottleneck 的conv2、conv3的 stride 参数。如果你觉得麻烦,resnet18/resnet34 这种基础版本不改变 stride,特征图 8×8 也勉强够用,但效果上限会低一些。
另外提一下预训练权重的加载。如果你用了单通道输入,conv1没法直接加载原版权重,怎么办?可以把原版conv1.weight在通道维度上求平均,得到一个新的 1×7×7 的卷积核。这样既利用了预训练的先验知识,又适配了单通道输入:
pretrained_conv1 = resnet.conv1.weight # shape: 64, 3, 7, 7 new_conv1_weight = pretrained_conv1.mean(dim=1, keepdim=True) # shape: 64, 1, 7, 7 model.conv1.weight.data = new_conv1_weight这个技巧很实用,强烈推荐。
3.3 Transformer 解码器的实现
解码器我直接用 PyTorch 内置的nn.TransformerDecoder和nn.TransformerDecoderLayer。这里选择标准 Transformer 而不是其他变种,原因很简单:内置模块稳定、文档多、不容易写出隐晦 bug,性能也足够。
参数上,d_model=512,nhead=8,num_layers=6,dim_feedforward=2048。虽然公式识别任务不需要超大模型,但 512 维是性能和显存的一个良好折中。如果显存紧张,可以降到 256,效果差距不会太大。
位置编码方面,我对比过两种:
- 固定正弦位置编码(原版 Transformer 用的)
- 可学习位置编码(每个位置分配一个可训练的向量)
在公式识别这个任务上,可学习位置编码效果稍好一些,因为公式图像的 token 长度相对固定(一般不超过 256),可学习编码可以针对实际长度做优化。实现时,nn.Embedding(max_len, d_model)就够了。
解码器的关键代码:
import torch import torch.nn as nn import math class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len=512): super().__init__() self.pe = nn.Embedding(max_len, d_model) def forward(self, x): # x: [seq_len, batch_size, d_model] seq_len = x.size(0) positions = torch.arange(seq_len, device=x.device) return x + self.pe(positions).unsqueeze(1) class TransformerDecoder(nn.Module): def __init__(self, vocab_size, d_model=512, nhead=8, num_layers=6, max_len=512): super().__init__() self.embedding = nn.Embedding(vocab_size, d_model) self.pos_encoder = PositionalEncoding(d_model, max_len) decoder_layer = nn.TransformerDecoderLayer( d_model=d_model, nhead=nhead, dim_feedforward=2048, batch_first=False, dropout=0.1 ) self.decoder = nn.TransformerDecoder(decoder_layer, num_layers=num_layers) self.fc_out = nn.Linear(d_model, vocab_size) def forward(self, tgt, memory, tgt_mask=None, tgt_key_padding_mask=None): tgt = self.embedding(tgt) * math.sqrt(self.embedding.embedding_dim) tgt = self.pos_encoder(tgt) output = self.decoder(tgt, memory, tgt_mask=tgt_mask, tgt_key_padding_mask=tgt_key_padding_mask) return self.fc_out(output)3.4 Mask 机制与推理时的 auto-regressive 逻辑
Transformer 解码器训练和推理有根本区别。训练时用的是 Teacher Forcing:一次把整个目标序列丢进去,通过 mask 保证位置 i 只能看到 i 之前的 token。推理时是自回归生成:先给<sos>,拿到第一个 token,再拼回去继续推理,直到遇到<eos>或者达到最大长度。
训练时的 mask 是两个:
tgt_mask:上三角掩码,让每个位置只能 attend 到左侧(包括自己)。tgt_key_padding_mask:<pad>位置的掩码,防止模型关注 padding token。
生成tgt_mask的代码:
def generate_square_subsequent_mask(sz): mask = torch.triu(torch.ones(sz, sz) * float('-inf'), diagonal=1) return mask这个 mask 会加到注意力分数上,让被 mask 位置的分数变成负无穷,softmax 之后权重就趋近于 0。很多初学者忘记做这个 mask,导致训练时模型“偷看未来”,训练 loss 很低但推理效果一塌糊涂,这里多检查几遍不亏。
4. 损失函数、训练策略与超参数调优
4.1 损失函数选型
序列生成任务最常用的损失函数是交叉熵损失。对于解码器输出的每个位置,我们都要计算预测 token 分布和真实 token 之间的交叉熵。但注意,<pad>位置需要被排除,不能参与 loss 计算。
PyTorch 里可以用nn.CrossEntropyLoss(ignore_index=pad_idx)来实现。这个ignore_index参数非常好用,不用手动做 mask。
我在训练中还使用了标签平滑(Label Smoothing)。标准交叉熵会鼓励模型对正确 token 给出接近 1 的置信度,容易导致过拟合。标签平滑把目标分布改成:正确 token 概率为 1-ε,其余 token 均匀分配 ε/(V-1)。这样做的好处是模型不会过度自信,对书写风格多变的手写体有更好的泛化能力。ε 我平时取 0.1。
4.2 Teacher Forcing 与 Scheduled Sampling
Teacher Forcing 是指训练时解码器的输入直接用真实标签序列,而不是模型自己的预测。这样做收敛很快,但会让训练和推理存在分布差异:训练时看到的是真实 token,推理时看到的却是自己生成的 token,一旦前面的 token 出错,错误会一路传播。
缓解这个问题的方式是 Scheduled Sampling:训练初期多用 Teacher Forcing,随着训练推进,以一定概率替换为模型自己的输出。这个概率可以按 epoch 递减,比如每个 epoch 增加 5% 的自生成比例。不过在公式识别这种任务上,我发现标准的 Teacher Forcing 配合 dropout 就已经够用了,Scheduled Sampling 对最终指标提升有限,还多了一堆超参要调,性价比不高。
4.3 优化器与学习率调度
优化器我选 AdamW,权重衰减设 1e-4。相比 Adam,AdamW 把权重衰减和梯度更新解耦,在 Transformer 这类模型上更稳定,不容易出现 loss 震荡。
学习率调度是整个训练策略里最讲究的一环。Transformer 对学习率非常敏感,直接用固定学习率很容易在初期发散。业界通用做法是“预热 + 衰减”:先让学习率从 0 线性升到峰值,再按余弦曲线慢慢降下来。PyTorch 里可以用get_cosine_schedule_with_warmup(来自 transformers 库)或者手写调度器。
我的配置是这样:
- 峰值学习率:1e-3
- warmup 步数:2000 步
- 总训练步数:约 50000 步
- 最小学习率:峰值学习率的 1/10
这个配置在 CROHME 数据集上训练大约 8 小时(单卡 RTX 3090)就能看到不错的效果。如果你显存小,可以调小 batch size,同时把学习率按比例降低,否则容易不稳定。
4.4 混合精度与梯度裁剪
如果显存紧张,强烈建议开启混合精度训练(AMP)。PyTorch 自带torch.cuda.amp,只需要改几行代码,显存占用能降低 30% 以上,训练速度也能提升不少。具体做法是:
scaler = torch.cuda.amp.GradScaler() for batch in dataloader: with torch.cuda.amp.autocast(): output = model(images, labels) loss = criterion(output, labels) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) scaler.step(optimizer) scaler.update()梯度裁剪 max_norm=1.0 对防止梯度爆炸非常关键。Transformer 训练时梯度范数偶尔会飙升,如果不裁剪,一个 batch 就能把模型参数冲飞,之前几个小时的训练全部白费。这个我经历过一次,损失从 0.5 直接变成 NaN,从头再来,非常痛苦。
4.5 训练超参数速查表
直接给一张我用到的超参配置表,方便大家直接抄作业:
| 参数 | 取值 | 说明 |
|---|---|---|
| 输入图像尺寸 | 256×256 | 长边缩放 + 白色 padding |
| 编码器 | ResNet18 | 可换成 ResNet34/50,注意显存 |
| d_model | 512 | Transformer 特征维度 |
| num_layers | 6 | 解码器层数 |
| nhead | 8 | 多头注意力头数 |
| dim_feedforward | 2048 | FFN 隐藏层维度 |
| batch_size | 32 | 单卡 24GB 显存可跑 |
| 优化器 | AdamW | lr=1e-3, weight_decay=1e-4 |
| 学习率调度 | warmup + cosine | warmup=2000 steps |
| label_smoothing | 0.1 | 降低过拟合 |
| dropout | 0.1 | 各子层 dropout |
| max_length | 256 | 输出序列最大长度 |
| beam_size | 5 | 推理时束搜索宽度 |
这个配置不一定是全局最优,但很稳。我试过把 d_model 提到 768、层数加到 8,效果确实提升,但训练时间几乎翻倍,对于课程项目来说性价比不高。
5. 推理与后处理
5.1 Greedy Search 与 Beam Search
推理时可以有多种解码策略。最简单的 Greedy Search 是每步选择概率最大的 token,直接作为下一步的输入,一直循环到<eos>。优点是速度快,缺点是容易陷入局部最优,一个 token 选错后面全崩。
Beam Search 是更稳妥的选择。它每一步保留概率最高的 K 个候选序列(K 叫 beam size),而不是只留一个。这样即使某个时间步的最佳候选最终证明是死路,还有另外 K-1 条路可以走。我用 beam size=5,效果比 Greedy 提升明显,尤其是在长公式上。
Beam Search 实现细节不算复杂,但要处理序列长度不一、结束条件不同步等问题。建议直接用开源库,比如torchaudio或ctcdecode不太适合这里,手写一个简洁版本更可控。核心逻辑就是维护一个候选列表,每步扩展,最后按得分排序。
5.2 长度惩罚
Beam Search 有一个天然问题:它倾向于短序列,因为概率是连乘的,序列越长概率越小。如果不加修正,模型会早早结束输出,漏掉后面的符号。所以需要一个长度惩罚系数,对长序列的得分做补偿。
常用公式是:
score = log_prob / (sequence_length ** length_penalty)
length_penalty 一般取 0.6~1.0。这个值是经验值,我通常先取 1.0 跑一版,再对比 0.6 的结果,选评测指标更好的那个。在公式识别任务上,1.0 的惩罚能让输出更完整,复现效果也更好。
5.3 无效 LaTeX 过滤与归一化
模型输出的序列不一定总是合法 LaTeX。比如可能出现未闭合的花括号、非法的上下标组合、重复的\frac导致嵌套过深等。这些情况需要在后处理阶段修正。
我的处理策略是先做括号匹配检查,把不成对的花括号补上或删掉。然后做简单的语法过滤,比如^和_后面必须跟一个合法 token,否则删除。这些规则虽然简单,但能显著提升最终渲染出的公式质量。
评测时还需要做归一化。CROHME 官方的做法是先把预测的 LaTeX 和真实 LaTeX 都转换成一个规范化的表示,比如去掉多余空格、统一\dfrac和\frac等,然后再计算准确率。如果不做归一化,很容易因为一个空格差异就被判错,白白丢分。
5.4 评估指标
公式识别任务的评估指标主要看两个:
- 表达式级准确率(Expression Accuracy):预测的 LaTeX 和真实 LaTeX 完全一致的比例。这个指标最严格,也最直观。
- Token 级准确率(Token Accuracy):预测序列和真实序列的 token 级匹配度,吃一点编辑距离的容错。通常用 BLEU 或者编辑距离来算。
做项目报告时,建议两个指标都汇报。表达式级准确率体现的是最终效果,token 级准确率能告诉你模型“差多少”,方便定位问题。在 CROHME 测试集上,ResNet + Transformer 这套方案通常能跑到 60%~70% 的表达式级准确率,对于课程项目来说已经是非常好的成绩了。
我用自己复现的模型跑了一版,测试集准确率在 65% 左右,主要错误集中在结构特别复杂的长公式上,例如多层嵌套的积分表达式。短公式和中等复杂度的式子表现很好,基本都能正确识别。
6. 常见问题与排查技巧
6.1 训练 Loss 不降或者下降极慢
遇到这种情况,先别急着加模型复杂度,按顺序排查:
- 数据顺序是否打乱:确认 DataLoader 的 shuffle=True,如果数据按公式类型排序,模型容易陷入局部最优。
- 学习率是否合适:把学习率调到 1e-3 左右(配合 warmup)再试。学习率过低,loss 下降会很慢。
- 标签序列是否正确:打印几个 batch 的 token id,人工确认一下标签有没有错位、缺失。这个问题看起来低级,但最容易发生。
- Mask 是否正确:检查 tgt_mask 是不是上三角,padding mask 有没有生效。Mask 错了模型能“偷看未来”,loss 前期会很低但验证集一塌糊涂。
6.2 训练集 Loss 低但验证集差
这是典型的过拟合。手写公式数据集小,模型很容易记住训练样本。优先做两件事:
- 加强数据增强,尤其是仿射变换和噪声。
- 增大 dropout,验证集效果不满意就把 dropout 从 0.1 提到 0.2 甚至 0.3。
另外,ResNet 预训练权重如果加载了,可以尝试冻结前几层不参与训练,只微调高层特征和 Transformer 解码器。这也能有效缓解过拟合,还能加快训练速度。
6.3 模型永远输出空序列
这种情况最让人抓狂:训练 loss 正常,但推理时模型只输出<sos>然后立刻输出<eos>,相当于一个字都没识别出来。
可能的原因有三个:
- 推理时没有用正确的方式生成,比如没有 mask 未来位置,导致模型每一步都看到“空”的后续区域。
- 解码器输入的 embedding 初始化有问题。
- 标签序列里
<sos>和<eos>的 id 搞反了。
我的排查方法是:先打印训练时 loss 是否收敛,再看一个 batch 的推理逐步输出。如果第二步就出现<eos>,大概率是<eos>在词表里的位置索引有问题,或者位置编码范围不对。
6.4 OOM(显存不足)
公式图像 padding 到 256×256,batch size 又大,显存确实压力不小。解决思路有这几个:
- 减小 batch size,同时按比例降低学习率。
- 开启混合精度训练,显存能降低很多。
- 图像尺寸降到 224×224(Imagenet 标准尺寸),效果损失不大。
- 用梯度累积,模拟更大的 batch。
梯度累积的实现很简单,每 N 个 batch 更新一次参数即可。比如实际 batch size 为 16,梯度累积 2 步,等效 batch size 为 32。
accumulation_steps = 2 optimizer.zero_grad() for i, batch in enumerate(dataloader): loss = model(**batch) loss = loss / accumulation_steps loss.backward() if (i + 1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad()6.5 生成的 LaTeX 渲染报错
模型推理出的字符串,如果直接丢给 LaTeX 编译器(比如渲染成图片再展示),经常会报“Undefined control sequence”之类的错误。这通常是因为输出里混入了词表中的奇怪 token,或者生成了非法的 LaTeX 命令。
我的做法是加一个白名单机制:推理时对不在 LaTeX 命令白名单里的 token 做过滤或替换。比如遇到\undeined这种非法命令,直接替换成空字符串。同时限制连续}的数量,避免嵌套结构崩坏。
7. 项目总结与扩展方向
写到这里,整个项目的核心内容基本讲完了。我个人在这个项目上最大的体会是:ResNet + Transformer 的组合看起来是“老技术拼凑”,但正是这种成熟技术的合理组合,解决了一个远比普通 OCR 复杂的问题。很多时候,做项目不需要追求“最新最潮的架构”,而是要把每个环节吃透,把细节做到位。数据预处理是否合理、Mask 有没有写对、学习率调度是否恰当,这些才是决定项目成败的关键。
这个项目后续还可以继续扩展的方向,我简单列几个:
- 把 ResNet 换成 Swin Transformer 或者 ConvNeXt,对比不同视觉编码器对公式识别效果的影响。
- 引入自监督预训练:先在大规模手写字符数据上做预训练,再在公式数据上微调,进一步提升效果。
- 将识别结果接入语音播报或者数学引擎(如 MathJax 渲染),做成一个完整的“手写公式拍照识别 + 计算”应用。
- 优化推理速度,部署到移动端或者 Web 端,这个方向对工程能力的要求更高,但项目含金量也更高。
最后再分享一个小技巧:调参时,每次只改一个变量,并且做好实验记录。我曾经为了赶时间一次性改了三个超参,结果模型崩了都不知道是哪个参数导致的。把每次实验的配置和指标记录下来,能帮你积累很多可复用的经验,这也是资深工程师和新手之间很明显的差距之一。
希望这篇分享能帮你把项目顺利做出来,并且真正理解背后每个环节的原理。有问题欢迎在评论区交流。
本文还有配套的精品资源,点击获取