☰
基于深度学习的端到端手写公式识别:从图像预处理到LaTeX推理全流程
2026/9/27 23:05:06 网站建设 项目流程

简介:一套基于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% 以下,几乎全栽在分割这一步。

端到端方案则避开了显式分割,用一个 CNN+RNN 的序列模型把整张图映射成 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, threshold=200) 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), interpolation=cv2.INTER_AREA) # 把图像填充到模型要求的 512x128,不足的部分补黑边 canvas = np.zeros((out_size[1], out_size[0]), dtype=np.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 定义第一版模型:一个能跑起来的 CNN+CTC 骨架

预处理做完后,下一步是定义模型。我习惯的配置是:CNN 特征提取 + 双向 GRU + CTC 解码。这个组合对公式这种“变长序列 + 二维布局”的任务比纯 CNN 好很多,而且参数量不大,CPU 也能训。

import torch import torch.nn as nn class FormulaNet(nn.Module): def __init__(self, num_classes, hidden_size=256): super().__init__() # CNN 部分把 512x128 的图压成特征序列,最后一维是时间步 self.cnn = nn.Sequential( nn.Conv2d(3, 32, kernel_size=3, padding=1), nn.BatchNorm2d(32), nn.ReLU(inplace=True), nn.MaxPool2d(2, 2), # 宽从 512 -> 256 nn.Conv2d(32, 64, kernel_size=3, padding=1), nn.BatchNorm2d(64), nn.ReLU(inplace=True), nn.MaxPool2d(2, 2), # 256 -> 128 nn.Conv2d(64, 128, kernel_size=3, padding=1), nn.BatchNorm2d(128), nn.ReLU(inplace=True), nn.MaxPool2d(2, 2), # 128 -> 64 nn.Conv2d(128, 256, kernel_size=3, padding=1), nn.BatchNorm2d(256), nn.ReLU(inplace=True), nn.MaxPool2d((2, 1)), # 高度压缩到 32,宽度保持 64 ) # 双向 GRU 建模字符之间的上下文关系 self.birnn = nn.GRU(256, hidden_size, bidirectional=True, batch_first=True) # 分类头:每个时间步输出一个字符概率分布 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_{i=1}^n i$ 要标成“\sum_{i=1}^{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_{i=1}^{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), dpi=128) fig.text(0.05, 0.4, f"${latex_str}$", fontsize=18) # 随机加旋转和缩放,模拟手写体常见的位移 ax = fig.gca() ax.set_axis_off() ax.set_xlim(0, 1); ax.set_ylim(0, 1) path = os.path.join(save_dir, f"img_{idx:05d}.png") fig.savefig(path, bbox_inches='tight', pad_inches=0.1) plt.close(fig) with open(os.path.join(save_dir, f"img_{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), dtype=torch.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-3

CTC 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_classes=len(token_to_idx)).to(device) model.load_state_dict(torch.load('checkpoints/model_best.pt', map_location=device)) 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.jpg

requirements.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', ascending=False).head(10))

验证时重点看 CER 最高的那批样本,它们通常集中在几个固定字符上,比如“l”与“1”、“\frac”和“\sqrt”的花括号配对错误。把这些案例用 matplotlib 和 OpenCV 画出来,其实就是一个典型的数据分析与可视化过程,能直观看到模型在哪些结构上最薄弱,我就从这里决定下一步的数据增强方向。

部署上,我会把训练好的 PyTorch 模型导出成 ONNX 格式,这样在本地 Python 环境里不需要依赖完整 torch 也能推理。导出命令和最小推理脚本如下:

import torch from model import FormulaNet model = FormulaNet(num_classes=len(token_to_idx)) model.load_state_dict(torch.load('checkpoints/model_best.pt', map_location='cpu')) 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 基线,否则旧模型很可能会和新字典错位,这几乎是所有翻车事故里最隐蔽的一类。希望这套从数据到验证的路径能帮你把先形成自己的手写公式识别闭环,希望帮到你。

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

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

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

立即咨询