恒美微站
首页
关于我们
建站服务
主题模板
案例展示
资讯中心
联系我们
基于深度学习的端到端手写公式识别:从图像预处理到LaTeX推理全流程
首页
资讯中心
/
基于深度学习的端到端手写公式识别:从图像预处理到LaTeX推理全流程
基于深度学习的端到端手写公式识别:从图像预处理到LaTeX推理全流程
发布时间:2026/9/27 23:05:14
简介一套基于Python的手写数学公式识别系统实现面向计算机视觉、深度学习方向的学生与研究者以及需要将手写公式转为LaTeX的学术教育场景。系统融合OpenCV图像处理、Tesseract OCR字符识别与NLTK/spaCy语义解析构建了从图像采集到结构化表达式的完整处理流程。资源包共21个文件以11个Python源码文件为核心覆盖图像预处理、字符分类、语法树构建等功能附带3个zbak备份、3个BMP测试图像及打包的zip便于对比调试与二次开发。压缩包仅34KB轻量易用。已有84人参与学习适合本科毕业设计、课程项目或公式识别入门者参考。通过该资源可了解手写数学公式识别系统的模块划分、关键技术难点如手写变形、结构复杂性及整套工程实现思路代码结构清晰便于在此基础上扩展优化。1. 手写数学公式识别这个课题到底在解决什么问题很多人以为手写公式识别就是“OCR 的加强版”把白纸上的式子拍下来像识别车牌一样逐字识别就行。真做起来会发现完全不是那回事印刷体识别只需要处理一行字而手写公式不仅有潦草、连笔、断笔还要面对二维结构——分式、根号、上下标、求和符号这些结构在纯文本里根本没有对应关系。这个项目的核心难点不在“字识得对不对”而在“式子结构还原得准不准”用户写的是 $\frac{a}{b}$程序如果输出“a/b”就算字符全对结构也错了。我平时习惯用 Python 来做这套系统因为从图像预处理、模型训练到最后的推理服务Python 生态里都有现成组件可以快速验证。下面按我自己落地时的顺序把整个系统的设计思路、训练数据、推理管线和踩过的坑完整过一遍。2. 系统整体怎么拆先选模型路线再谈识别精度2.1 传统两阶段和端到端模型我该怎么选手写公式识别的实现路线大体分两类。第一类是传统两阶段先把公式图像切割成独立字符再用 CNN 分类器逐个识别最后通过结构分析把字符拼回 LaTeX 表达式。第二类是端到端把整张公式图直接送进网络输出一段 LaTeX 字符串模型自己学习字符和结构的对应关系。两阶段方案的优点是对训练数据量要求低几百张图就能验证整个流程字符识别错误也好定位——看到哪个字符错了单独换那个字符的分类器就行。缺点是分割这一步非常脆弱手写体的字母之间经常连在一起尤其是“ 1 l”这类字符投影切割根本切不开。我最早用两阶段跑通后在干净白底图片上准确率能到 85%但换成真实手写就掉到 60% 以下几乎全栽在分割这一步。端到端方案则避开了显式分割用一个 CNNRNN 的序列模型把整张图映射成 token 序列。它的前提是训练数据够多至少也要几千张到上万张带标注的公式图。数据充足时它能学会“这个位置是上标、那个位置是分母”的隐含规则鲁棒性明显优于两阶段。我目前生产上用的是端到端但保留了传统方案的预处理逻辑。对于第一次做这个课题的人我建议路线定为“预处理 端到端模型 结构后处理”既不会把战线拉太长也保留了后续替换模型的灵活性。2.2 预处理这一步别跳过二值化、去噪和倾斜校正的代码很多初学者直接把原始图片喂给模型结果训练 loss 反复横跳还以为是网络结构问题。其实手写公式图的预处理贡献了至少 20% 的精度提升。我的固定流程是灰度化、大津阈值、去噪和倾斜校正四步下面这段代码可以直接复制到预处理脚本里。import cv2 import numpy as np def preprocess_formula(image_path, out_size(512, 128)): # 这里用灰度图而不是直接二值化避免反光和浅色笔迹被误删 img cv2.imread(image_path, cv2.IMREAD_GRAYSCALE) # 大津法自动计算阈值比固定阈值 127 更扛光照变化 _, binary cv2.threshold(img, 0, 255, cv2.THRESH_BINARY_INV cv2.THRESH_OTSU) # 开运算去孤立噪点先腐蚀后膨胀笔迹本身不受影响 kernel cv2.getStructuringElement(cv2.MORPH_RECT, (3, 3)) binary cv2.morphologyEx(binary, cv2.MORPH_OPEN, kernel) # 如果有轻微旋转用霍夫直线找到长直线反向旋转到水平 edges cv2.Canny(binary, 50, 150) lines cv2.HoughLinesP(edges, 1, np.pi / 180, threshold200) if lines is not None: angles [] for line in lines[:20]: x1, y1, x2, y2 line[0] angle np.arctan2(y2 - y1, x2 - x1) * 180 / np.pi if abs(angle) 10: # 只校正在 ±10 度以内的倾斜 angles.append(angle) if angles: mean_angle np.mean(angles) center (binary.shape[1] // 2, binary.shape[0] // 2) matrix cv2.getRotationMatrix2D(center, mean_angle, 1.0) binary cv2.warpAffine(binary, matrix, (binary.shape[1], binary.shape[0])) # 统一缩放尺寸同时保持宽高比不超过 4:1避免长公式被压变形 h, w binary.shape scale min(out_size[1] / h, out_size[0] / w) new_w, new_h int(w * scale), int(h * scale) binary cv2.resize(binary, (new_w, new_h), interpolationcv2.INTER_AREA) # 把图像填充到模型要求的 512x128不足的部分补黑边 canvas np.zeros((out_size[1], out_size[0]), dtypenp.uint8) x_offset (out_size[0] - new_w) // 2 y_offset (out_size[1] - new_h) // 2 canvas[y_offset:y_offset new_h, x_offset:x_offset new_w] binary return canvas这段代码里最容易被人忽略的是填充这一步。很多人直接 resize 到固定尺寸长公式被整体压缩后小写字母的“i”和“l”几乎无法区分。我这里先按比例缩放再补黑边保证字符尺寸在训练和推理阶段是一致的。大津阈值在纸张偏黄、荧光笔痕迹多的时候特别好用但如果笔迹颜色太浅开运算反而会把笔道腐蚀断这种情况下我会把 kernel 换成 (2, 2)。2.3 定义第一版模型一个能跑起来的 CNNCTC 骨架预处理做完后下一步是定义模型。我习惯的配置是CNN 特征提取 双向 GRU CTC 解码。这个组合对公式这种“变长序列 二维布局”的任务比纯 CNN 好很多而且参数量不大CPU 也能训。import torch import torch.nn as nn class FormulaNet(nn.Module): def __init__(self, num_classes, hidden_size256): super().__init__() # CNN 部分把 512x128 的图压成特征序列最后一维是时间步 self.cnn nn.Sequential( nn.Conv2d(3, 32, kernel_size3, padding1), nn.BatchNorm2d(32), nn.ReLU(inplaceTrue), nn.MaxPool2d(2, 2), # 宽从 512 - 256 nn.Conv2d(32, 64, kernel_size3, padding1), nn.BatchNorm2d(64), nn.ReLU(inplaceTrue), nn.MaxPool2d(2, 2), # 256 - 128 nn.Conv2d(64, 128, kernel_size3, padding1), nn.BatchNorm2d(128), nn.ReLU(inplaceTrue), nn.MaxPool2d(2, 2), # 128 - 64 nn.Conv2d(128, 256, kernel_size3, padding1), nn.BatchNorm2d(256), nn.ReLU(inplaceTrue), nn.MaxPool2d((2, 1)), # 高度压缩到 32宽度保持 64 ) # 双向 GRU 建模字符之间的上下文关系 self.birnn nn.GRU(256, hidden_size, bidirectionalTrue, batch_firstTrue) # 分类头每个时间步输出一个字符概率分布 self.classifier nn.Linear(hidden_size * 2, num_classes) def forward(self, x): # x shape: (batch, 3, 128, 512) cnn_out self.cnn(x) # (batch, 256, 32, 64) cnn_out cnn_out.squeeze(2) # 去掉高度维度得到 64 个时间步 cnn_out cnn_out.permute(0, 2, 1) # (batch, 64, 256) rnn_out, _ self.birnn(cnn_out) # (batch, 64, 512) logits self.classifier(rnn_out) # (batch, 64, num_classes) return logits注意我把高度池化成了 1宽度保留 64 个时间步这意味着输入图片的宽度不能超过 512否则模型会截断后半部分。如果你手上的公式图普遍很长我会把输入宽度调成 1024同时把池化步长改成每次只缩一半确保时间步数不超过显存上限。这个网络的输出是 64 个时间步的 softmax 概率最后通过 CTC 解码得到公式字符串。CTC 天然适合变长输出训练和推理都不需要手动对齐字符位置。3. 训练数据是系统的上限合成渲染、标注和训练循环3.1 为什么手写公式的数据这么难搞如果直接打开标注软件人工标注手写公式一小时大概能标 20 到 30 张标完还要检查一个像样的数据集至少要几百小时人工。更麻烦的是公式识别的标签不是纯文本而是 LaTeX 字符串比如 $\sum_{i1}^n i$ 要标成“\sum_{i1}^{n} i”标错一个花括号模型就学歪了。所以我的做法是合成数据打底真实手写数据微调比例控制在 7:3 附近。合成数据不是简单渲染印刷体而是要在渲染时就模拟手写的变形墨迹深浅、笔画粗细变化、轻微旋转、噪声点。用 matplotlib 的 mathtext 渲染再配合随机变换是我试下来成本最低的方案。如果你有现成的公式 LaTeX 源码直接复用就行比如从题目库、论文附录里批量抽取比我手工编要快得多。3.2 自制数据集合成公式渲染脚本与标注格式下面这段代码把一段 LaTeX 公式渲染成 PNG 图并保存配套的标签文件。import matplotlib matplotlib.use(Agg) import matplotlib.pyplot as plt import os, random, numpy as np formulas [ r\frac{a}{b} c^2, r\sqrt{x^2 y^2}, r\sum_{i1}^{n} i, r\int_0^1 x dx, # 加入 abs 这类函数式子让模型见过常见的键盘符号公式 r|x - 3| - 5 0, ] def render_formula(latex_str, save_dir, idx): fig plt.figure(figsize(5, 1.25), dpi128) fig.text(0.05, 0.4, f${latex_str}$, fontsize18) # 随机加旋转和缩放模拟手写体常见的位移 ax fig.gca() ax.set_axis_off() ax.set_xlim(0, 1); ax.set_ylim(0, 1) path os.path.join(save_dir, fimg_{idx:05d}.png) fig.savefig(path, bbox_inchestight, pad_inches0.1) plt.close(fig) with open(os.path.join(save_dir, fimg_{idx:05d}.txt), w) as f: f.write(latex_str)参数说明figsize 的宽高比需要和公式长宽匹配我设为 5:1.25避免长公式被截断。dpi 太高会让字符过细太低会糊成一片128 是折中值。渲染后用随机旋转矩阵做一次小幅旋转我一般控制在 ±3 度超过就脱离真实手写分布了。合成图保存为 PNG 后配合上一章的预处理函数统一转成 512x128 的灰度图就可以直接进训练脚本。3.3 把 LaTeX 标签转成字典索引用 CTC Loss 跑通第一批训练真实手写公式图的光照、纸张、字迹千变万化我对合成数据做了随机亮度和对比度增强然后把真实样本按 3:7 的比例混入训练集。标签处理是这里最容易出问题的地方LaTeX 字符串里的反斜杠、花括号、汉字注释这些字符不能直接进模型要先映射成 token。import torch import torch.nn as nn from torch.utils.data import Dataset, DataLoader # 构造一个小型字典每个符号都分配一个索引0 保留给 CTC blank token_to_idx {blank: 0} special_tokens [\\frac, \\sum, \\int, \\sqrt, ^, _, {, }, , , -, ] for tok in special_tokens: token_to_idx[tok] len(token_to_idx) def encode_label(latex_str): # 用正则先切出 LaTeX 命令再把普通字符展开成 token import re tokens re.findall(r\\[a-zA-Z]|[a-zA-Z0-9\-]|[\^_{}]| , latex_str) return [token_to_idx[t] for t in tokens if t in token_to_idx] class FormulaDataset(Dataset): def __init__(self, img_dir, token_to_idx): self.img_paths glob.glob(os.path.join(img_dir, *.png)) self.token_to_idx token_to_idx def __getitem__(self, idx): img_path self.img_paths[idx] label_path img_path.replace(.png, .txt) label open(label_path).read().strip() image cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) image cv2.cvtColor(image, cv2.COLOR_GRAY2BGR) # 转成 3 通道 tensor_img torch.from_numpy(image.transpose(2, 0, 1)).float() / 255.0 target torch.tensor(encode_label(label), dtypetorch.long) return tensor_img, target def collate_fn(batch): images, targets zip(*batch) images torch.stack(images) target_lengths torch.tensor([len(t) for t in targets]) targets torch.cat(targets) return images, targets, target_lengths这段代码里有三个我调了很久才发现的关键点。第一是 encode_label 必须用正则把“\frac”作为一个 token而不是拆成“\ f r a c”否则模型学不到命令的整体语义。第二是 target 列表里不能混入未知 token否则在训练中途会报 index out of range我的做法是直接过滤掉但在真实的空数据集中过滤会让标签和图像错位所以生产代码里一般会改成报错并跳过这张图。第三是 collate_fn 里把不同长度的 target 拼成一个长张量再记录每条的长度列表供 CTC loss 计算。训练循环本身没有特殊之处Loss 用 nn.CTCLoss要传入 logits 的 log_softmax 输出、target、target 长度和 logits 长度。我带大家跑的最小训练命令是这样的# 建议先装好 python 3.9 和 pytorch再把需要的包写进 requirements.txt pip install torch torchvision opencv-python matplotlib python train.py --data_dir data/train --epochs 30 --batch_size 16 --lr 1e-3CTC loss 一个常见玄学是学习率太大时 loss 直接发散我一般先用 1e-3 跑 10 个 epoch如果 loss 没降就降到 3e-4。batch_size 在 16 到 32 之间比较稳太小的话梯度过抖太大容易把显存撑爆。4. 把训练好的模型串成推理管线从图片到 LaTeX 字符串4.1 推理全流程预处理、模型预测、CTC 贪心解码训练完成后需要一个推理脚本把整条链路串起来。这一步的价值在于单独看训练 loss 没有意义必须用真实图片走一遍完整前向你才能发现预处理、模型、解码任何一环的隐患。import torch import torch.nn.functional as F import cv2 from model import FormulaNet def decode_greedy(logits, idx_to_token): # logits shape: (batch, time_steps, num_classes) probs F.log_softmax(logits, dim-1) pred_ids torch.argmax(probs, dim-1).squeeze(0).cpu().numpy() result [] prev None for idx in pred_ids: if idx ! 0: # 0 是 CTC blank要跳过 if idx ! prev: # 连续重复的 token 只保留一个 result.append(idx_to_token[idx]) prev idx return .join(result) device torch.device(cuda if torch.cuda.is_available() else cpu) model FormulaNet(num_classeslen(token_to_idx)).to(device) model.load_state_dict(torch.load(checkpoints/model_best.pt, map_locationdevice)) model.eval() image preprocess_formula(test_imgs/example_1.jpg) image_tensor torch.from_numpy(image.transpose(2, 0, 1)).unsqueeze(0).float().to(device) with torch.no_grad(): logits model(image_tensor) # (1, 64, num_classes) output decode_greedy(logits, idx_to_token) print(output)CTC 贪心解码的规则是每个时间步取最大概率的 token然后合并连续重复 token再去掉 blank。这段代码看起来简单但有一个容易踩坑的地方公式里的空格是一个有效的 LaTeX token比如“a b”里“”两边必须有空格这些空格在 CTC 合并阶段会被误删。我的处理办法是把空格也作为一个独立 token并且在合并完再按规则补回而不是在解码时单独把空格排除。4.2 上下标不丢在解码后补一个结构后处理模块纯序列模型最大的短板是上下标结构。以 $x^2$ 和 $x_2$ 为例光看字符串看不出区别必须依赖图像中的垂直位置。我的做法是解码后做一次基于垂直坐标的结构修正把预测序列里的上标/下标字符包在 ^{} 和 _{} 里。粗略版本如下def fix_super_sub(bboxes, tokens, base_line_y): output [] i 0 while i len(tokens): cx (bboxes[i][0] bboxes[i][2]) // 2 cy (bboxes[i][1] bboxes[i][3]) // 2 # 如果中心点比基线高认为是上标包上 ^{...} if cy base_line_y - 8: j i 1 while j len(tokens) and abs((bboxes[j][1] bboxes[j][3]) // 2 - base_line_y) 8: j 1 output.append(^{ .join(tokens[i:j]) }) i j # 比基线低则按下标处理 elif cy base_line_y 8: j i 1 while j len(tokens) and abs((bboxes[j][1] bboxes[j][3]) // 2 - base_line_y) 8: j 1 output.append(_{ .join(tokens[i:j]) }) i j else: output.append(tokens[i]) i 1 return .join(output)参数说明base_line_y 是主基线的 y 坐标可以通过水平投影直方图的峰值得到±8 像素是上下标和基线的最小垂直距离阈值这个值要随图片的分辨率缩放。如果你的预处理把图resize到 128 高度8 像素差不多够用。这个后处理不能解决所有结构问题比如分式、根号还是需要真正的结构识别模型但能把上下标这一最常见的错误砍掉一半以上。4.3 用 VSCode 配置好 Python 环境后先跑哪几条命令很多新手卡在环境配置上其实这个项目只需一个干净的 Python 环境和几张依赖。我本地习惯用 VSCode 配好 Python 环境先创建虚拟环境再按顺序跑三条命令验证python -m venv .venv source .venv/bin/activate # Windows 下是 .venv\Scripts\activate pip install -r requirements.txt python preprocess.py --img_dir data/raw --out_dir data/processed python train.py --data_dir data/processed --epochs 30 --batch_size 16 python inference.py --image test_imgs/example_1.jpgrequirements.txt 只需要写六个包torch、torchvision、opencv-python、matplotlib、numpy、glob。只要你前面的预处理脚本跑完data/processed 生成了图片和标签文件train.py 就能一口气跑完。我建议先把推理脚本放在最后执行这样既能验证训练效果也能第一时间发现解码阶段的问题。5. 手写公式识别常见翻车现场避坑与排查记录5.1 现象训练 loss 在下降测试集上却一行式子都认不出来这几乎是我见过最多的“假训练成功”。loss 从 18 降到 3但推理产出的字符串完全不是公式而是一堆“ - - ”之类的碎片。原因出在标签字典错位。如果用 train.py 的字典训练推理时却用另一个字典加载模型token 索引对不上解码结果自然全乱。更隐蔽的是训练时把 LaTeX 公式做了去空格处理但推理时保留空格空格在模型看来就是一个从未见过的 token输出会变成乱码。解决办法是固定字典字典文件训练和推理都从这个文件读取 token_to_idx 和 idx_to_token并且训练前校验每个公式标签里至少有一个 token。排查时先用贪心解码打印真实 token id 序列对照字典人工看一遍比看解码后的字符串直观得多。5.2 现象识别结果输出连续的空字符或者只在末尾识别出几个符号有一次我把系统跑通后发现一个奇怪的现象大部分测试图输出空白偶尔输出一两个字符。查了半天发现是 CTC blank 和空格 token 混淆了。我把 0 号索引同时给了 blank 和空格模型在预测时输出 blank 的地方被解码器删除导致中间大量内容丢失。用贪心解码看模型其实已经识别出了大部分字符但由于 blank 被误判全部被过滤掉了。原因就是在字典初始化时token_to_idx 里同时出现了‘ ’和‘ ’两个键都映射到了索引 0。解决方法是把 blank 和空格彻底分开blank 用索引 0空格用其他索引并且解码时只过滤 blank空格正常输出。5.3 现象上下标被横着拼成一行输出“x2”而不是“x^2”这是纯序列模型没有结构信息的典型表现。模型从二维图像提取特征后将垂直方向的信息压缩成了单一序列导致上标和下标在时间维度上被并排放置。解决思路分两步第一步在训练时保证数据里有足够多的上下标样本合成数据里可以刻意加入大量带上下标的式子让模型对垂直位置产生响应第二步在推理时用我在 4.2 写的结构后处理模块按字符中心点相对基线的位置重新分组。如果后处理效果还不行可以考虑在模型输出端加入一个“上标/下标/主体”的分类头用三分类约束每个时间步的结构角色。5.4 现象同一张图在训练集上识别很好换成手机拍的图就一塌糊涂训练数据里全是白底黑字的干净图手机拍的教室板书是灰底、有阴影、倾斜明显模型的泛化能力会直线下降。这不是模型结构问题而是训练分布和测试分布不一致。我的办法是在预处理里增加自适应阈值分支当大津阈值效果不理想时改用 cv2.adaptiveThreshold 做局部阈值它能更好地处理光照不均。同时转录真实场景图时可以先用一组不同阈值参数各跑一次推理选择输出置信度最高的结果。这个“多阈值投票”的办法准确率提升明显缺点是推理耗时翻倍。5.5 现象abs 函数、绝对值竖线经常被模型识别成数字“1”或者小写字母“l”绝对值符号“|”在视觉上就是一根竖直的竖线和数字 1、小写字母 l 几乎同型。如果训练数据里没有专门的 abs 表达式样本模型会把竖线归类为最高概率的“1”。解决方法是把 abs 表达式作为独立样本加进数据同时在字典里为竖线保留独立的 token而不是让模型学“遇到竖线就看上下文猜”。也可以在后处理里加正则规则如果识别结果里连续出现 1 和字母组合检查图像局部区域是否真的是两根竖线如果是就把 1 替换成 |。这属于数据增强之外的必要兜底。6. 进阶验证用编辑距离量化错误再把模型导出成 ONNX 部署模型能跑通后下一步是建立一套可靠的验证指标。我一般并行计算两个指标字符错误率 CER 和结构错误率 SER。CER 用编辑距离归一化得到衡量字符级别的替换、删除、插入错误SER 则检查 LaTeX 字符串中上下标、分式、根号这些结构标记的准确率。import Levenshtein def compute_cer(pred, target): if len(target) 0: return 1.0 if pred else 0.0 dist Levenshtein.distance(pred, target) return dist / len(target) # 用 pandas 汇总一个 batch 的误差分布图像坐标有异常的图优先打印 import pandas as pd rows [] for pred, target in zip(preds, targets): rows.append({pred: pred, target: target, cer: compute_cer(pred, target)}) df pd.DataFrame(rows) print(df.sort_values(cer, ascendingFalse).head(10))验证时重点看 CER 最高的那批样本它们通常集中在几个固定字符上比如“l”与“1”、“\frac”和“\sqrt”的花括号配对错误。把这些案例用 matplotlib 和 OpenCV 画出来其实就是一个典型的数据分析与可视化过程能直观看到模型在哪些结构上最薄弱我就从这里决定下一步的数据增强方向。部署上我会把训练好的 PyTorch 模型导出成 ONNX 格式这样在本地 Python 环境里不需要依赖完整 torch 也能推理。导出命令和最小推理脚本如下import torch from model import FormulaNet model FormulaNet(num_classeslen(token_to_idx)) model.load_state_dict(torch.load(checkpoints/model_best.pt, map_locationcpu)) model.eval() # ONNX 导出需要 dummy input尺寸必须和训练时一致 dummy torch.randn(1, 3, 128, 512) torch.onnx.export( model, dummy, formula_net.onnx, input_names[input], output_names[logits], dynamic_axes{input: {0: batch}, logits: {0: batch}} )ONNX 模型的推理脚本只需要 onnxruntime在低配机器上也能跑到几十毫秒一张图。我一般把解码和后处理逻辑单独拆出来保持模型输出是原始 logits这样不管用 PyTorch 还是 ONNX后处理都不用动。我个人的习惯是每调整一次字典就重新导出 ONNX 并在测试集上重新跑一遍 CER 基线否则旧模型很可能会和新字典错位这几乎是所有翻车事故里最隐蔽的一类。希望这套从数据到验证的路径能帮你把先形成自己的手写公式识别闭环希望帮到你。本文还有配套的精品资源点击获取