2026/9/10 11:17:50

端到端公式识别:从CNN+Transformer到LaTeX生成实战

端到端公式识别:从CNN+Transformer到LaTeX生成实战 简介这是一套基于 PyTorch 的端到端图像 LaTeX 公式识别项目面向有一定基础的深度学习开发者与科研教育人员解决从数学公式图像到可编辑 LaTeX 代码自动转换的难题。内容覆盖图像预处理、卷积神经网络特征提取、循环神经网络序列解码以及编码器-解码器结构配有训练、验证、测试所需的 json 标注文件、词汇表与 7 个可运行 Python 脚本适合希望复现公式识别流程并加深对关键环节理解的读者。压缩包共有 122 个文件以 110 张公式样本图像为主另含 4 个 json 配置文件和 1 份 README 说明总体积只有 378KB轻量易用。目录结构清晰、分模块组织README 中附环境搭建说明便于按需查阅。目前已有 138 人学习浏览项目源码结合环境搭建指引可帮助快速上手完整走通数据准备、模型训练到性能评估的全流程。1. 为什么公式识别最适合用 End-to-End 的方式做题库系统、论文库、在线教育内容加工这些场景里文字 OCR 早就不是瓶颈真正麻烦的是公式区——普通 OCR 对着一串积分和矩阵输出的要么是乱码要么是一张无法检索的图片。公式识别的难点不在认符号而在猜结构分式的分子分母归属、求和符号上下限的层级、花括号配对全依赖版面二维布局。传统做法是「符号检测 符号分类 结构分析」逐段拼装步骤一多错误率就叠加上去。End-to-End 的思路是把「图像像素 → LaTeX 源码文本」直接建成序列到序列模型编码器吃灰度图解码器逐 token 产出 LaTeX 代码整体是一个模型训练和部署都在 PyTorch 框架里闭环。对做论文入库、错题录入、LaTeX 公式检索的团队来说这是当前维护成本最低的路线。下面按架构选型、数据准备、训练配置、推理调优这条落地路径展开。2. 从图像到 LaTeX token为什么编码器和解码器都要为版面服务2.1 两阶段方案的问题结构错误在每一步累积早年业界常用的 pipeline 包含三个独立阶段——符号检测、符号分类、基于规则或图模型的结构分析。符号检测阶段遇到根号横杠、求和符号上下限、绝对值竖线时召回率本身就不高到了结构分析阶段上下标归属和分数线配对又依赖前一步的分类置信度任何一步出错都会被放大。这类方法在干净白底黑字的合成数据上表现尚可一遇到真实扫描件里公式半角全角混杂、字号不一调试成本会直接超过重写一个模型。2.2 端到端模型怎么看到二维结构端到端方案里图像在多个尺度上被编码成特征图。解码器生成\frac时注意力会落在分式横杠附近的特征区域生成上标^时注意力落在右上或右下的小字区域。也就是说二维排版的隐式建模由注意力机制承担不需要显式规则表。训练拉普拉斯公式\frac{d^2 y}{dx^2}这样的样本时解码器的注意力确实会在分子分母之间来回切换可视化时非常直观。这里有一个常见误区把图片直接 resize 成正方形再喂给模型。公式图像是长条形正方形化会把字号压扁或拉长导致\partial这类带弧线的符号严重变形。常见做法是统一把高度缩放到 64 或 96 像素宽度按比例缩放超出定宽的做右侧 padding不足的也补齐既不破坏字符比例也保住 GPU 上的 batch 形状。2.3 CNN 特征 2D 位置编码是稳妥组合端到端公式识别的主流实现是 CNN 编码器加 Transformer 解码器CNN 用 ResNet-18/34 或 DenseNet 提取视觉特征Transformer 负责把特征序列解码成 LaTeX token。近两年也出现 ViT 编码器方案在跨行分数、矩阵这类长距离结构建模上更占优但数据量和训练轮次要翻倍。两边对比编码器方案优势代价ResNet 可学习 2D 位置编码收敛快、对低分辨率输入友好长距离依赖需要堆更多层ViT 正弦 2D 位置编码全局看得全、矩阵和超大括号生成稳定训练数据不够时方差大这一列在实际项目里的决策含义是手头只有几万张合成公式图就从 ResNet 起步数据量到了几十万量级再切 ViT 不迟。做图像算法选型时不要先追新架构而是先看手里数据的规模能不能喂饱它。2.4 训练目标不是整句匹配而是逐 token 交叉熵PyTorch 里实现损失非常直接解码器每个位置输出一个 logits 向量与真实 token 对齐后计算 CrossEntropyLoss序列整体取平均。这里的 Pascalignore_index必须指向 padding token否则未对齐位置会把梯度带偏。label smoothing 也是必备项数值调到 0.1 附近具体原因在 4.3 里展开。3. 数据生成比模型调参更关键LaTeX 渲染管线与字符集设计3.1 用本机 TeX 直接渲染训练样本公开的公式识别数据集规模有限常见做法是自己搭一条渲染管线随机组合公式模板再调用 TeX 引擎渲染成图片。需要注意不要用 matplotlib 的 mathtext 兜底它对\begin{aligned}、\begin{matrix}支持很差而这两个命令在高年级题库里高频出现。下面这段代码是渲染管线的最小实现# generate_pairs.py import subprocess from pathlib import Path out_dir Path(formula_pairs) out_dir.mkdir(exist_okTrue) def render_formula(tex_source: str, out_prefix: str, dpi: int 150): workdir out_dir / out_prefix workdir.mkdir(exist_okTrue) tex_path workdir / formula.tex # standalone 类会按公式内容自动裁剪边界去掉多余留白 tex_path.write_text( r\documentclass[border1pt]{standalone} \n r\usepackage{amsmath,amssymb} \n r\begin{document} \n tex_source \n r\end{document} ) # dvipng 比 pdftocrop 少一层转换生成速度更快 subprocess.run([latex, -interactionnonstopmode, formula.tex], cwdworkdir, checkFalse, capture_outputTrue) subprocess.run([dvipng, -D, str(dpi), -T, tight, formula.dvi, -o, f{out_prefix}.png], cwdworkdir, checkFalse, capture_outputTrue) if (workdir / f{out_prefix}.png).exists(): # 文本配对一行 LaTeX 源码对应一行图片路径 with open(out_dir / train_pairs.txt, a, encodingutf-8) as f: f.write(f{workdir / (out_prefix .png)}\t{tex_source}\n)逻辑说明standalone文档类负责裁剪白边dvipng -T tight进一步压缩图片尺寸。train_pairs.txt保存图片绝对路径和 LaTeX 源码训练脚本按行读取即可。dpi参数想模拟真实扫描件可以降到 100120想模拟高清拍照可以拉高到 300 再随机 resize。为了覆盖不同字体来源渲染时把 Computer Modern 和 STIX 字体各跑一遍能显著减少l、1、\ell之间的混淆。3.2 字符集设计的两个不变量「LaTeX 符号大全」看着很长实际训练时不能贪多。字符集是公式识别最容易被低估的环节我坚持两个原则第一等价命令归一化。\dfrac一律归一成\frac\left(直接输出(所有\displaystyle直接剔除。模型输出里的模板噪声越少结构错误越少。第二基础 token 分两层覆盖。一层是abcdefghijklmnopqrstuvwxyz、数字、常用标点、-*/()另一层是\frac \sqrt \sum \int \prod \partial \mathrm \mathbf \begin{matrix}这类结构命令。希腊字母只收高频项\varepsilon,\varphi一类等题库分布确定后再决定是否加入。用脚本检查语料里出现但 vocab 缺失的命令比人工对着符号大全核对靠谱得多import re missing set(re.findall(r\\[a-zA-Z], all_tex_sources)) - set(vocab) print(missing) # 把漏网之鱼直接揪出来3.3 数据增强按破坏可读性来选公式识别里几何变换要克制。旋转角度一般不超过 5 度超过 15 度模型就会开始乱透视变换适合模拟拍照角度但变化太大会让分数线弯曲这与真实扫描件的形变方向不一致。我一般固定用四个增强随机 0.81.2 倍缩放、亮度抖动、高斯噪声、小范围裁剪。模拟模糊时可以借用图像超分辨率重建里的降采样思路——先缩小到 0.5 倍再放大回原尺寸比单纯加高斯模糊更贴近真实低清输入尤其适合处理翻拍课本的公式。4. 用 PyTorch 搭一个能跑的最小训练闭环4.1 编码器与解码器的骨架代码下面代码只保留主干backbone 换成自己的数据路径就能当 baseline# model.py import torch import torch.nn as nn from torchvision.models import resnet18 class FormulaEncoder(nn.Module): def __init__(self, enc_dim256): super().__init__() backbone resnet18(weightsNone) # 去掉最后的全局池化和分类头只留卷积特征 self.features nn.Sequential(*list(backbone.children())[:-2]) self.proj nn.Conv2d(512, enc_dim, 1) # 可学习 2D 位置编码高度 4、宽度 16对应下采样 16 倍 self.pos_h nn.Parameter(torch.randn(4, enc_dim)) self.pos_w nn.Parameter(torch.randn(16, enc_dim)) def forward(self, x): feat self.features(x) # b, 512, h, w feat self.proj(feat) # b, 256, h, w b, c, h, w feat.shape pos (self.pos_h[:h].unsqueeze(1) self.pos_w[:w].unsqueeze(0)) feat feat pos.permute(2, 0, 1).unsqueeze(0) # 拼成序列先列后行方便解码器按行扫描公式 return feat.flatten(2).permute(2, 0, 1) # s, b, c class FormulaDecoder(nn.Module): def __init__(self, vocab_size, d_model256, nhead8, num_layers4): super().__init__() self.embed nn.Embedding(vocab_size, d_model) self.pos nn.Parameter(torch.randn(192, d_model)) layer nn.TransformerDecoderLayer(d_modeld_model, nheadnhead, batch_firstFalse, dropout0.1) self.decoder nn.TransformerDecoder(layer, num_layersnum_layers) self.out_proj nn.Linear(d_model, vocab_size) def forward(self, tgt, memory, tgt_mask): # tgt 在训练时已经是 [BOS] tokens[:-1] 的拼接 tgt self.embed(tgt) self.pos[:tgt.size(0)] out self.decoder(tgt, memory, tgt_masktgt_mask) return self.out_proj(out)逻辑说明resnet18下采样 16 倍64×256 的输入变成 4×16 的特征图2D 位置编码逐元素加到特征图上解码器才能分辨「同行不同列」和「同列不同行」。memory是编码器输出的全部特征序列Transformer 解码器每个时刻通过交叉注意力从中取用与当前 token 相关的版面位置。参数说明d_model256在公式识别里够用nhead8兼顾多头注意力的拆分num_layers4是训练速度和结构建模能力的平衡点层数上到 6 在 5 万张以内的合成数据上容易过拟合。vocab_size取决于字符集设计一般控制在 10003000不宜再大。4.2 训练循环里的 teacher forcing 与 mask公式解码必须用带因果掩码的逐 token 生成PyTorch 的TransformerDecoderLayer需要外部传入tgt_maskdef build_causal_mask(length): return torch.triu(torch.full((length, length), float(-inf)), diagonal1)训练时真实 token 序列整体送入解码器tgt_mask保证第 i 个位置的注意力只看得到前 i 个 token。这个 mask 不加模型准确率会直接掉到只能输出第一个字符。4.3 损失、优化器与单卡训练配方损失用交叉熵时把ignore_index设为pad_idx避免 padding 位置参与梯度计算。优化器选 AdamW权重衰减 0.01学习率峰值 1e-4前 1000 步 warmup 线性上升到峰值之后按步数余弦退火。参数项建议值原因图像高度64 或 96 px低于 48 时小字号上下标粘连严重序列最大长度192128 对长公式截断会导致括号缺失batch size16 梯度累积到 32省显存且不牺牲批次多样性label smoothing0.1让括号类强配对 token 不被过度压制学习率1e-4偏保守长任务更稳定验证时最常用的两个指标是 Exact Match 和编辑距离。EM 对公式识别偏严格一个多余的\left就判错编辑距离更符合线下评测场景——用户能看懂、能编译的公式就应该算有效输出。我只用 EM 当门禁指标模型在第 20 个 epoch 前后 EM 的收益明显放缓紧接着就要去检查编辑距离的分布看错是错在符号级还是结构级。5. 推理解码与后处理Beam search 宽度要保守括号校验必须有5.1 一个够用的 Beam Search 片段公式识别的推理阶段Greedy 解码容易在\frac和}之间漏掉子结构。常见做法是 beam search 取 5 条候选保留得分最高且能通过后处理校验的序列。def beam_step(logits, scores, seqs, vocab_size, beam_width): log_probs torch.log_softmax(logits[:, -1], dim-1) next_scores scores.unsqueeze(1) log_probs flat_scores next_scores.view(-1) topk torch.topk(flat_scores, beam_width) parent_ids topk.indices // vocab_size token_ids topk.indices % vocab_size seqs torch.cat([seqs[parent_ids], token_ids.unsqueeze(1)], dim1) return seqs, topk.values逻辑说明beam search 不是每个时刻独立取 top-k而是把上一步的 beam 宽度乘以词表大小后统一排序再保留全局前 k 条路径。参数建议beam_width 取 5 比较合适取 10 以上不仅慢一倍错误率未必下降——公式的歧义空间比自然语言小宽束容易把低分路径也带上来。5.2 后处理三步模型产生的{、}、\left、\right一旦不配对输出 LaTeX 无法编译伤害远大于错别字。我通常在后处理里做三件事检查花括号配对只保留能配平的候选不配平就走下一条 beam而不是强行补}。把\left和\right.这类容易遗漏的转义 token 单独回填注意不要改动括号之间的内容。命令归一化\dfrac替换为\frac\displaystyle删除\top按题库偏好归一为^T。5.3 高频失败场景与对策现象根因对策长公式尾部丢\right)解码长度被 max_len 截断把 max_len 提到 256并先对公式区域做切分合成数据 EM 高、真实图片 EM 掉一半真实场景带下划线或荧光笔噪声混合「划痕线」与「色块涂抹」两类输入显存溢出feature map 宽度过大限制图像宽度 512 以内超出的分区识别最后说一个我一直在用的验证技巧让模型同时输出 token 序列和最后一层的 token 置信度部署时凡是置信度低于 0.7 的序列先在服务端用本机 LaTeX 编译一次能编译过再返回给用户。编译不过的候选直接降级把对应图片路径和错误 LaTeX 写进难例日志每周把这些样本挑出来加进下一轮训练集。这个「能编译才算通过」的口径比任何编辑距离阈值都更贴近用户真实使用场景也是公式识别任务里一道廉价兜底闸门。本文还有配套的精品资源点击获取