☰
ResNet+Transformer:端到端手写数学公式识别流水线解析
2026/10/2 3:21:05 网站建设 项目流程

简介:一套面向深度学习、计算机视觉课程设计及毕业设计场景的手写数学公式识别Python源码,采用ResNet提取视觉特征、Transformer编码器-解码器完成序列生成,实现端到端识别。项目按工程化方式组织,配置管理、数据加载、词表构建、模型编解码、训练评估与测试推理等模块划分清楚,适合有一定深度学习基础、需要完成高分开题/结题项目或复现该方向算法的学习者。压缩包共32个文件,以19个Python源文件为主,附带对应pyc缓存、2个txt字典与配置说明、1个cfg工程配置、1个YAML配置及1个打包结果,整体仅87KB,轻量易部署。代码经过严格调试,可在本地环境直接运行,并附有单一识别结果输出,便于快速验证效果。目前已有646人学习,可参考其数据组织、词表构建与端到端训练流程,适合作为课程大作业、竞赛方案或算法入门基线。

1. 手写数学公式识别:把 ResNet 和 Transformer 拧成一条识别流水线

期末大作业要交一个「能跑、能讲、能答辩」的项目,手写数学公式识别是个很讨巧的题目:有视觉难度,有序列生成难度,技术栈还正好踩在 CNN 和 Transformer 的交界处。这份源码的核心思路是端到端:输入一张手写公式图片,输出对应的 LaTeX 序列。视觉部分用 ResNet 提取特征,序列部分用 Transformer 解码器自回归出 token,中间用注意力机制把图像特征和文字序列对齐。它解决的是「图片里的积分号、根号、分式怎么变成一段能编译的 LaTeX 代码」这个具体问题,适合正在做课程设计、准备毕设、或者想自己跑通一个 OCR 序列生成项目的开发者。资源拿到手不是用来 Readme 考古的,是让你把训练、推理、调参、答辩这一整条链路走完的。

2. 从源码结构到模型原理:先搞清楚每个文件是干嘛的

2.1 解压之后先别跑:把目录结构和数据流摸清楚

拿到源码包第一件事不是pip install -r requirements.txt,而是打开目录结构看数据流。这类项目文件通常不多,但每个文件都卡在一条流水线上:图片进、LaTeX 出。我习惯先画一条数据流:dataset.py负责把图片和标注读成 batch,models/里放着 ResNet 编码器和 Transformer 解码器,train.py把两者串起来做前向和反向,predict.py负责推理时把概率序列转成可读文本。

├── config.py # 全局配置:图像尺寸、token数、学习率、beam size ├── dataset.py # 数据加载:图片预处理 + LaTeX 标注转 token ├── models │ ├── resnet.py # 视觉编码器,输出特征图序列 │ ├── transformer.py # 解码器:自注意力 + 交叉注意力 │ └── attention.py # 注意力封装,含 masked self-attention ├── train.py # 训练入口 ├── predict.py # 单张图片推理入口 ├── utils │ ├── tokenizer.py # LaTeX 字符串 <-> token id │ └── metrics.py # 字符准确率、序列准确率 ├── checkpoints/ # 模型权重存放处 └── data/ ├── images/ # 公式图片 └── labels.txt # 每行:图片名 + LaTeX 标注

这种结构是典型的「编码器-解码器」框架。resnet.py里不是直接把 ResNet 最后池化层的输出拿来用,而是把 feature map 保留成C x H x W的形状,再按空间位置展开成(H*W) x C的序列——这一步是 CNN 和 Transformer 之间的桥梁。transformer.py里的解码器接收这个序列作为交叉注意力的 Key 和 Value,同时用自回归方式逐个预测 token。数据流的顺序决定了你改代码的顺序:先改数据集,再改模型,最后动训练脚本。

2.2 ResNet 做视觉特征、Transformer 做序列生成:为什么是这对组合

这套架构不是拍脑袋凑的。手写公式识别比普通 OCR 难在两点:一是符号多且形近(\sum和\Sigma只差一个笔画粗细),二是结构复杂(分式、上下标、根号嵌套)。纯 CNN 可以把图像分类做得很好,但输出不定长序列需要序列模型;纯 RNN 能生成序列,但对长距离的视觉依赖关系建模弱。Transformer 的自注意力恰好能把「积分号在图片左上角、被积函数在右边」这种跨区域关系一次性建模,而 ResNet 的强特征提取能力保证了输入给注意力的是高质量的表征。

具体到实现上,常见做法是用 ResNet-18 或 ResNet-50 做骨干网络,去掉最后的全连接层和全局池化,保留C x H x W的特征图。然后做一个维度对齐:因为 Transformer 的d_model通常是 256 或 512,而 ResNet 最后一层输出通道是 512 或 2048,中间要加一个1x1卷积把通道数压到d_model。之后再按空间位置展开成序列,加上位置编码。位置编码我建议先试可学习的nn.Embedding,因为公式图片里符号的位置偏移比自然场景更敏感,可学习编码在训练数据量不大的时候比正弦编码更容易收敛。

Transformer 解码器的输入是目标序列(LaTeX token),通过 masked self-attention 保证每个位置只能看到当前位置之前的 token。每一层解码器还有一个交叉注意力子层,Query 来自解码器自身,Key 和 Value 来自 ResNet 输出的视觉特征序列。这个交叉注意力就是模型「看图片写字」的核心机制——你在推理时可以把注意力权重可视化出来,能看到生成\frac的时候模型确实在关注图片里分数线附近的位置。

2.3 训练流程拆解:数据怎么进、损失怎么算、梯度怎么回传

训练循环本身不复杂,但有几个细节决定了能不能收敛。数据进模型之前,图片会被缩放到固定尺寸(常见是64 x 256或48 x 192),但同时要做宽高比保持和 padding,否则公式里的长根号、长分数线会被压变形,这属于细粒度特征的破坏。LaTeX 标注先通过 tokenizer 转成 id 序列,序列头加<sos>尾加<eos>,padding 到固定长度,padding 位置在损失计算时要 mask 掉。

# train.py 训练循环核心片段 for batch in dataloader: images, labels, label_len = batch # images: (B, 3, H, W) -> ResNet -> (B, H*W, d_model) features = encoder(images) # labels: (B, seq_len),前向时用 teacher forcing logits = decoder(features, labels[:, :-1]) # 输入不含 eos # logits: (B, seq_len-1, vocab_size) loss = criterion( logits.reshape(-1, vocab_size), labels[:, 1:].reshape(-1) # 预测目标是从第二个 token 开始 ) # mask 掉 padding 位置 mask = labels[:, 1:] != pad_idx loss = (loss * mask).sum() / mask.sum() optimizer.zero_grad() loss.backward() clip_grad_norm_(model.parameters(), 5.0) # 防止梯度爆炸 optimizer.step()

这段代码里最需要注意的是labels[:, :-1]和labels[:, 1:]的错位:模型输入是「前 N-1 个 token」,预测目标是「后 N-1 个 token」,这样每个位置都在预测下一个 token。mask的计算很关键,padding 位置不计入损失,否则模型会花大量精力去学「预测 pad 本身」,拉低真实符号的收敛速度。梯度裁剪的阈值设在5.0,Transformer 在训练初期特别容易梯度异常,这个操作基本是标配。

损失函数用nn.CrossEntropyLoss(ignore_index=pad_idx)也行,但我更推荐手动 mask,原因在避坑章节会展开解释。学习率一般初始1e-4到3e-4,配合 warmup 策略:前 5 个 epoch 线性升到峰值,之后按步长衰减。如果你发现 loss 下降得特别慢,先检查是不是学习率太小,再看 teacher forcing 有没有写对——这两个是最常见的「不收敛」来源。

3. 环境准备与数据组织:跑起来之前先解决两个实际问题

3.1 依赖安装:PyTorch 版本和 CUDA 的匹配

源码包的requirements.txt一般会写torch>=1.9,但实际踩坑的人都知道,光有这个远远不够。PyTorch 的安装方式直接决定后面所有环节顺不顺畅。我建议用 conda 建独立环境,Python 版本选 3.8 或 3.9,不要追新——有些老代码在 Python 3.11 下会因为torchvision的 API 变更直接报错。

conda create -n formula python=3.9 conda activate formula # CUDA 11.8 对应 PyTorch 2.0+,CUDA 12.1 对应 PyTorch 2.1+ pip install torch torchvision --index-url https://download.pytorch.org/whl/cu118 pip install -r requirements.txt

装完先跑一个快速验证,确认 CUDA 真的能用,别等到训练跑了一半才发现用的是 CPU:

python -c "import torch; print(torch.cuda.is_available(), torch.cuda.get_device_name(0))"

如果输出True和显卡型号,环境这关算过了。requirements.txt里通常还有numpy、opencv-python、pillow、tqdm,这些版本要求不高,直接装最新问题不大。但注意一点:如果项目里用了torchtext,那就得小心版本了——torchtext0.9 到 0.12 的 API 变化比较大,老代码里torchtext.data.Field这种写法在新版本里已经被删了。遇到这种情况,最简单的处理是把数据加载逻辑改成torch.utils.data.Dataset实现,别在 torchtext 上死磕。

3.2 数据格式与标注组织:label.txt 就是你的训练契约

数据是这类项目能不能复现的最大变量。源码通常会附带一个示例数据集,但规模很小(几十到几百张),只够跑通流程。如果你想训练出一个真正能用的模型,需要自己扩充数据。数据格式一般是labels.txt,每行两列:图片文件名 + LaTeX 标注。

img_001.png \frac { a } { b } + \sqrt { x } img_002.png \sum _ { i = 1 } ^ { n } i ^ { 2 }

注意标注里的空格不是随意的——tokenizer 依赖空格进行分词。\frac { a } { b }会被拆成['\frac', '{', 'a', '}', '{', 'b', '}']这样一个个 token,如果你写成\frac{a}{b}不带空格,分词结果就是['\frac{a}{b}']一个整体,模型根本学不到分子分母的结构。这一点在检查别人数据集的时候要格外留意。

图片预处理在dataset.py里做,典型流程是:读图 -> 灰度化 -> 二值化 -> 保持宽高比缩放到目标尺寸 -> padding 到统一大小。很多初学者直接cv2.resize(img, (W, H))硬压,结果把\lim里的点和\cdot压成噪点。正确做法是先按长边缩放,再在短边补零。代码里对应的是:

# dataset.py 中图片预处理的核心逻辑 def preprocess_image(img_path, target_w=256, target_h=64): img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) h, w = img.shape scale = min(target_h / h, target_w / w) new_w, new_h = int(w * scale), int(h * scale) img = cv2.resize(img, (new_w, new_h), interpolation=cv2.INTER_AREA) canvas = np.ones((target_h, target_w), dtype=np.uint8) * 255 x_offset = (target_w - new_w) // 2 y_offset = (target_h - new_h) // 2 canvas[y_offset:y_offset + new_h, x_offset:x_offset + new_w] = img return canvas

这里scale = min(target_h / h, target_w / w)是核心:取两个缩放比例里较小的那个,保证图片完整放进目标画布,不裁剪任何符号。INTER_AREA插值在缩小图片时能保留边缘信息,比INTER_LINEAR更适合二值化后的手写笔画。padding 用 255(白色背景)而不是 0(黑色),因为手写公式通常画在白纸上,保持背景一致性可以减少模型对背景色的过拟合。

3.3 tokenizer 的完整流程:从 LaTeX 字符串到 token id 序列

Tokenizer 是这类项目里最容易被低估的组件。它做的事情看起来简单——字符串转 id、id 转字符串——但边界情况特别多:\frac这种带反斜杠的命令、{}这种结构符号、^_这种上下标标记、还有\alpha\beta这种希腊字母。一个健壮的 tokenizer 需要维护一张「token 到 id」的映射表,同时处理未登录词(OOV)。

# utils/tokenizer.py 的核心逻辑 class FormulaTokenizer: def __init__(self, vocab_file): self.token2idx = {} self.idx2token = {} with open(vocab_file, 'r', encoding='utf-8') as f: for line in f: token = line.strip() idx = len(self.token2idx) self.token2idx[token] = idx self.idx2token[idx] = token self.pad_idx = self.token2idx['<pad>'] self.sos_idx = self.token2idx['<sos>'] self.eos_idx = self.token2idx['<eos>'] def encode(self, latex_str): # 按空格分词,加入起止符 tokens = latex_str.strip().split(' ') ids = [self.sos_idx] + [self.token2idx[t] for t in tokens if t in self.token2idx] + [self.eos_idx] return ids def decode(self, ids): # 跳过特殊符,拼接回 LaTeX tokens = [self.idx2token[i] for i in ids if i not in (self.pad_idx, self.sos_idx, self.eos_idx)] return ' '.join(tokens)

第一次写encode时容易漏掉<sos>和<eos>。如果没有起止符,训练时模型不知道从哪里开始生成、在哪里结束,推理时就会一直生成下去直到撞上最大长度限制。decode时把特殊 token 过滤掉,避免输出里出现<eos>这种字符串。vocab 文件由build_vocab.py之类的脚本从训练标注里统计生成,一般做法是统计所有出现过的 token,过滤掉出现次数少于 2 的稀有 token。

4. 训练与推理实战:命令行参数、日志观察和单张图片测试

4.1 训练命令和关键参数:从默认配置到手动调优

源码包里通常有一个默认的config.py,但直接跑默认参数大概率不是最优的。先看一眼关键配置长什么样:

# config.py 中的核心配置 class Config: # 数据 data_dir = './data' batch_size = 16 shuffle = True num_workers = 4 # 模型 backbone = 'resnet18' # resnet18 / resnet34 / resnet50 d_model = 256 # Transformer 特征维度 nhead = 8 # 多头注意力头数 num_encoder_layers = 0 # 部分实现不用编码器,直接交叉注意力 num_decoder_layers = 4 # 解码器层数 max_len = 150 # 最大序列长度 # 训练 epochs = 60 lr = 2e-4 warmup_steps = 1000 grad_clip = 5.0 log_interval = 20 # 推理 beam_size = 5 length_penalty = 1.0

d_model=256、nhead=8、num_decoder_layers=4这套组合在公式识别场景下是性价比比较高的配置。d_model再大的话参数量涨得很快,但收益有限,因为手写公式的「词汇量」通常只有两三百个 token,不像机器翻译需要很大的模型容量。num_encoder_layers=0值得注意:很多实现里 ResNet 本身就充当了编码器角色,不需要再叠加 Transformer 编码器层,直接把 CNN 特征送进解码器的交叉注意力。如果你看到源码里encoder和decoder都有,那就是把 ResNet 特征又过了几层 Transformer 编码器,理论上能建模更全局的视觉依赖,但训练时间也会变长。

训练命令一般是:

python train.py --config config.py --data_dir ./data --epochs 60 --batch_size 16

训练日志会输出每个 batch 的 loss、当前学习率、以及每跑完一个 epoch 在验证集上的字符准确率。我习惯盯两个指标:第一个是 loss 是否在稳步下降,如果前 5 个 epoch 内 loss 纹丝不动,大概率是学习率或数据有问题;第二个是验证集字符准确率是否在 20 个 epoch 左右开始超过 50%,如果迟迟上不去,检查交叉注意力是不是没有正确接收 ResNet 的特征。

4.2 推理流程:beam search 生成 LaTeX 序列

推理阶段不再使用 teacher forcing,而是让模型自回归生成。常见做法是用 beam search 保留多个候选序列,避免贪心解码掉进局部最优。推理代码长这样:

# predict.py 的推理流程 def predict(model, img_tensor, tokenizer, config): model.eval() with torch.no_grad(): features = model.encoder(img_tensor) # (1, H*W, d_model) # beam search 初始化 beams = [([tokenizer.sos_idx], 0.0)] # (token序列, 累积log概率) for step in range(config.max_len): new_beams = [] for seq, score in beams: if seq[-1] == tokenizer.eos_idx: new_beams.append((seq, score)) continue seq_tensor = torch.tensor(seq).unsqueeze(0) logits = model.decoder(features, seq_tensor) # (1, step+1, vocab) next_log_probs = logits[0, -1, :].log_softmax(-1) top_k = next_log_probs.topk(config.beam_size) for token_id, log_prob in zip(top_k.indices, top_k.values): new_beams.append((seq + [token_id.item()], score + log_prob.item())) # 保留 beam_size 个最优 new_beams.sort(key=lambda x: x[1] / len(x[0]) ** config.length_penalty, reverse=True) beams = new_beams[:config.beam_size] if all(seq[-1] == tokenizer.eos_idx for seq, _ in beams): break best_seq = beams[0][0] return tokenizer.decode(best_seq)

beam search 的几个细节直接影响输出质量。length_penalty是对长序列的惩罚系数:设得太小,模型倾向于输出短序列,复杂的根号分式结构会被截断;设得太大,可能输出冗余的 token。常见起步值是1.0。max_len设 150 是给复杂的多重嵌套公式留余量,但如果你的数据集中公式都很短,可以降到 100,减少无效计算。topk 的beam_size=5是准确率和速度的折中,beam_size=10会好一点但推理时间翻倍。

推理时还要注意一个细节:在seq[-1] == tokenizer.eos_idx时直接保留该 beam,不再扩展。如果不做这个判断,beam 会在生成<eos>之后继续生成无意义的 token,白白浪费计算量。整个算法跑完之后,还需要做一个后处理:把 token 序列里的特殊符过滤掉,再把空格拼回去,得到最终的 LaTeX 字符串。

4.3 验证指标:字符准确率和序列准确率怎么算

训练完不能只看 loss,要跑验证集算指标。公式识别有两个常用指标:字符准确率(Character Accuracy)和序列准确率(Sequence Accuracy)。字符准确率是逐 token 比较预测和真值,允许部分正确;序列准确率要求整条 LaTeX 序列完全一致,更严格。实现时用编辑距离和精确匹配:

# utils/metrics.py 中的指标计算 def compute_metrics(pred_ids, true_ids, pad_idx): # 去掉 padding 和特殊符 pred = [i for i in pred_ids if i not in (pad_idx, sos_idx, eos_idx)] true = [i for i in true_ids if i not in (pad_idx, sos_idx, eos_idx)] char_acc = sum(1 for p, t in zip(pred, true) if p == t) / max(len(true), 1) seq_acc = 1.0 if pred == true else 0.0 return char_acc, seq_acc

字符准确率高但序列准确率低的场景很典型:根号里面的内容识别对了,但\frac和\sqrt的嵌套顺序反了。遇到这种情况不要急着调模型,先看具体错误样本,判断是注意力对齐问题还是数据标注问题。

5. 避坑指南:训练和推理中常见的五个翻车现场

5.1 现象:forward 时报维度不匹配,size mismatch

报错信息通常是The size of tensor a (512) must match the size of tensor b (256)。原因基本可以锁定在 ResNet 输出通道数和 Transformer 的d_model不一致。ResNet-18 最后一层输出是 512 通道,ResNet-50 是 2048,而d_model可能配的是 256。解决方式是在 ResNet 后加一个1x1卷积或线性层做投影。

# models/resnet.py 中加一个投影层 self.proj = nn.Conv2d(512, d_model, kernel_size=1) # 512 -> 256 def forward(self, x): features = self.backbone(x) # (B, 512, H', W') features = self.proj(features) # (B, d_model, H', W') B, C, H, W = features.shape return features.flatten(2).permute(0, 2, 1) # (B, H*W, d_model)

这个proj层只做通道数对齐,不做空间变换,所以用kernel_size=1就够了,不会破坏特征图的空间结构。从那以后我每次拿到新配置第一件事就是核对这两个数字。

5.2 现象:训练 loss 下降很慢,十几个 epoch 还在 5 以上

原因有几种:学习率太小、标签序列没有加起止符、或者 padding 部分参与了损失计算。其中 padding 参与损失是最隐蔽的——如果CrossEntropyLoss没有设ignore_index=pad_idx,模型会花费大量容量去学习预测 padding token,而对真实公式符号的学习被稀释了。解决方式在前面已经提过,训练循环里手动计算 mask 后再算 loss,或者直接用nn.CrossEntropyLoss(ignore_index=pad_idx)。

学习率方面,如果初始lr=2e-4在 5 个 epoch 内看不到明显下降,可以试着调到5e-4,但要同步检查是否有 loss 震荡——震荡说明已经过大了。加入 warmup 策略也能缓解前期不稳定的问题。

5.3 现象:推理时生成了重复 token,比如\frac \frac \frac

原因多半是 beam search 的停止条件没写好。如果没有在生成<eos>时及时终止,模型会在结束符之后继续预测,产生大量重复内容。另一个可能是模型训练数据里没有见过<eos>,导致推理时模型不知道何时停止。解决方式是在encode时确保每条标注都加了<eos>,并且在 beam search 中一旦某个 beam 生成了<eos>,就把它移到完成列表,不再扩展。

5.4 现象:小符号识别率特别低,比如\cdot、\prime经常被漏掉

这类问题通常出在图片预处理环节。如果直接把图片硬缩放到64x256,\cdot这种细粒度符号可能只有几个像素宽,特征完全丢失。解决方式是保持宽高比缩放到更高分辨率(比如80x320),配合 padding 而不是裁剪。另外,ResNet 的高层特征图空间分辨率低,小符号的信息到深层已经没了,可以尝试用 FPN 结构把浅层和深层特征融合,这一点在下一章展开。

5.5 现象:GPU 显存溢出,CUDA out of memory

最常见的原因是batch_size太大,或者序列长度过长导致注意力矩阵过大。Transformer 的注意力复杂度是O(n^2),公式图片展开成序列后长度可能到几百,加上 beam search 的多个候选同时喂进模型,显存很容易爆。解决方式是减小batch_size,用梯度累积保持等效 batch,或者开启混合精度训练:

python train.py --batch_size 4 --grad_accum 4 --fp16

如果显存只有 4GB,把d_model从 256 降到 192 也是一个选项,效果差一些但能跑起来。也可以用torch.utils.checkpoint对 ResNet 部分做梯度检查点,用计算换显存。

6. 进阶技巧:FPN 融合粗粒度与细粒度特征,再加上注意力可视化调试

6.1 用 FPN 结构融合多层 ResNet 特征

基础版 ResNet 编码器只用最后一层特征图,而最后一层分辨率低,丢失了细粒度信息。手写公式里的\cdot、\prime、\ddot都依赖高分辨率特征。在处理这类「小目标」问题时,FPN(Feature Pyramid Network)是常用解法:把 ResNet 的 C2、C3、C4、C5 四层特征图逐级上采样并相加,让大符号和小符号的信息都能保留在最终特征里。

# models/resnet.py 中 FPN 特征融合的简洁实现 class ResNetFPN(nn.Module): def __init__(self, d_model): super().__init__() self.backbone = models.resnet18(pretrained=True) # 取出各层特征 self.layer1 = self.backbone.layer1 # C2, 64通道 self.layer2 = self.backbone.layer2 # C3, 128通道 self.layer3 = self.backbone.layer3 # C4, 256通道 self.layer4 = self.backbone.layer4 # C5, 512通道 # 各层对齐到 d_model self.proj2 = nn.Conv2d(64, d_model, 1) self.proj3 = nn.Conv2d(128, d_model, 1) self.proj4 = nn.Conv2d(256, d_model, 1) self.proj5 = nn.Conv2d(512, d_model, 1) def forward(self, x): c2 = self.layer1(x) # (B, 64, H/4, W/4) c3 = self.layer2(c2) # (B, 128, H/8, W/8) c4 = self.layer3(c3) # (B, 256, H/16, W/16) c5 = self.layer4(c4) # (B, 512, H/32, W/32) p5 = self.proj5(c5) p4 = self.proj4(c4) + F.interpolate(p5, size=c4.shape[-2:], mode='bilinear') p3 = self.proj3(c3) + F.interpolate(p4, size=c3.shape[-2:], mode='bilinear') p2 = self.proj2(c2) + F.interpolate(p3, size=c2.shape[-2:], mode='bilinear') # 融合后展平。空间分辨率不同,可以取最大层做序列长度基准 return p2.flatten(2).permute(0, 2, 1) # (B, H/4 * W/4, d_model)

这里的核心思想是「高层特征有语义、低层特征有细节」,相加融合后给到 Transformer 的序列既包含整体结构信息(根号、分式的大框架),又保留笔画级细节(点、撇、小符号)。引入 FPN 后序列长度会变大(因为用的是 C2 层的分辨率),训练和推理会变慢,但准确率通常能涨 3 到 5 个百分点。

6.2 可视化交叉注意力权重,定位识别错误来源

训练完之后如果某些公式总是识别错,不要只盯着 loss 看,把交叉注意力权重画出来。具体做法是:推理时把 decoder 最后一层每个 head 的注意力矩阵导出来,对生成公式的某个 token(比如\frac),看它在生成时对图片上哪些位置关注最多,保存成灰度热力图。

# 可视化注意力权重(简化版) def dump_attention(model, img_tensor, seq_ids, save_path): attn_weights = [] def hook_fn(module, input, output): attn_weights.append(output[1].detach().cpu()) # 交叉注意力权重 handle = model.decoder.cross_attention.register_forward_hook(hook_fn) _ = model(img_tensor.unsqueeze(0), seq_ids.unsqueeze(0)) handle.remove() # attn_weights: (layers, B, heads, tgt_len, src_len) attn = attn_weights[-1][0, 0] # 只看最后一个注意力层、第一个head plt.imshow(attn, cmap='hot', aspect='auto') plt.xlabel('Image features'); plt.ylabel('Generated tokens') plt.savefig(save_path)

实际操作中我会挑一个识别失败的样本,逐 token 查看它生成时的注意力分布。如果模型生成\sqrt时注意力集中在图片的右上角,说明它关注的区域不对,问题可能出在特征提取阶段;如果注意力分布很分散,说明模型没有找到对应的视觉锚点,可以考虑增加训练数据里类似样本的比例。热力图比 loss 曲线直观得多,能少走很多弯路。

6.3 推理时的 LaTeX 语法后处理

模型输出的是 token 序列,拼成字符串后未必是合法 LaTeX,常见的错误包括\frac缺参数、花括号不匹配。一个简单的后处理是补全花括号:

def postprocess(latex_str): # 去掉多余空格,统一格式 latex_str = latex_str.replace(' { ', '{').replace(' } ', '}') # 确保花括号闭合 open_braces = latex_str.count('{') close_braces = latex_str.count('}') if open_braces > close_braces: latex_str += '}' * (open_braces - close_braces) return latex_str

这个后处理能在评测时把序列准确率拉高几个点,因为很多「错误」其实只是格式问题——\frac { a } { b }和\frac{a}{b}在数学上是同一个公式,但在文本匹配时会被计为错误。从那以后我每次跑完推理都会强制走一遍后处理,再上指标,这个习惯让最终报告里的数字好看不少,希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询