简介:这是一套基于 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_ok=True) def render_formula(tex_source: str, out_prefix: str, dpi: int = 150): workdir = out_dir / out_prefix workdir.mkdir(exist_ok=True) tex_path = workdir / "formula.tex" # standalone 类会按公式内容自动裁剪边界,去掉多余留白 tex_path.write_text( r"\documentclass[border=1pt]{standalone}" + "\n" r"\usepackage{amsmath,amssymb}" + "\n" r"\begin{document}" + "\n" + tex_source + "\n" r"\end{document}" ) # dvipng 比 pdftocrop 少一层转换,生成速度更快 subprocess.run(["latex", "-interaction=nonstopmode", "formula.tex"], cwd=workdir, check=False, capture_output=True) subprocess.run(["dvipng", "-D", str(dpi), "-T", "tight", "formula.dvi", "-o", f"{out_prefix}.png"], cwd=workdir, check=False, capture_output=True) if (workdir / f"{out_prefix}.png").exists(): # 文本配对:一行 LaTeX 源码对应一行图片路径 with open(out_dir / "train_pairs.txt", "a", encoding="utf-8") as f: f.write(f"{workdir / (out_prefix + '.png')}\t{tex_source}\n")逻辑说明:standalone文档类负责裁剪白边,dvipng -T tight进一步压缩图片尺寸。train_pairs.txt保存图片绝对路径和 LaTeX 源码,训练脚本按行读取即可。dpi参数想模拟真实扫描件可以降到 100~120,想模拟高清拍照可以拉高到 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.8~1.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_dim=256): super().__init__() backbone = resnet18(weights=None) # 去掉最后的全局池化和分类头,只留卷积特征 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_model=256, nhead=8, num_layers=4): super().__init__() self.embed = nn.Embedding(vocab_size, d_model) self.pos = nn.Parameter(torch.randn(192, d_model)) layer = nn.TransformerDecoderLayer(d_model=d_model, nhead=nhead, batch_first=False, dropout=0.1) self.decoder = nn.TransformerDecoder(layer, num_layers=num_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_mask=tgt_mask) return self.out_proj(out)逻辑说明:resnet18下采样 16 倍,64×256 的输入变成 4×16 的特征图;2D 位置编码逐元素加到特征图上,解码器才能分辨「同行不同列」和「同列不同行」。memory是编码器输出的全部特征序列,Transformer 解码器每个时刻通过交叉注意力从中取用与当前 token 相关的版面位置。
参数说明:d_model=256在公式识别里够用,nhead=8兼顾多头注意力的拆分;num_layers=4是训练速度和结构建模能力的平衡点,层数上到 6 在 5 万张以内的合成数据上容易过拟合。vocab_size取决于字符集设计,一般控制在 1000~3000,不宜再大。
4.2 训练循环里的 teacher forcing 与 mask
公式解码必须用带因果掩码的逐 token 生成,PyTorch 的TransformerDecoderLayer需要外部传入tgt_mask:
def build_causal_mask(length): return torch.triu(torch.full((length, length), float("-inf")), diagonal=1)训练时真实 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 时小字号上下标粘连严重 |
| 序列最大长度 | 192 | 128 对长公式截断会导致括号缺失 |
| batch size | 16 + 梯度累积到 32 | 省显存且不牺牲批次多样性 |
| label smoothing | 0.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)], dim=1) 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 写进难例日志,每周把这些样本挑出来加进下一轮训练集。这个「能编译才算通过」的口径,比任何编辑距离阈值都更贴近用户真实使用场景,也是公式识别任务里一道廉价兜底闸门。
本文还有配套的精品资源,点击获取