简介:这份资源是一套基于深度学习ResNet架构的手写数学公式识别系统,面向教育领域的技术开发者、机器学习学习者及中小学数学教学场景,用于自动识别手写公式中的加减乘除、括号及0-9数字并完成计算。系统采用Python 3.9开发,涵盖数据集预处理、模型训练与推理、图形界面交互等完整流程,ResNet的残差结构有效缓解深层网络梯度消失问题,提升复杂图像特征提取能力。压缩包共43个文件,约39.85MB,包含13个py源码文件、20张png图像、1个pth模型权重及txt、md、docx等说明文档,覆盖训练脚本、分割模块、预测逻辑与UI组件,目录结构清晰便于按模块学习。目前已有69人学习下载。读者可获取从数据清洗、归一化到模型部署的完整工程实现,理解手写公式识别与计算的技术链路,并借助说明文档快速上手运行与二次开发,适合作为教育智能化方向的实践参考。
1. 手写公式识别到底难在哪:从一张草稿纸说起
教育场景里有个很具体的痛点:学生在平板上写完(3+5)×2-4÷2,系统只能存成一张图片,没法判断对错,更没法自动批改。手写数学公式识别要解决的就是这件事——把笔迹图像转成结构化的 LaTeX 或表达式树,再交给计算引擎求值。它和普通 OCR 最大的区别在于:数字和运算符的识别只是第一步,真正的难点是二维结构。x²和x2像素分布接近,1/2的横线位置决定它是分数还是两个独立数字,括号的跨度可能横跨三行。ResNet 在这里的价值是提供一个足够强的视觉骨干,把笔画特征抽干净,后面再接序列解码或结构分析模块。这套方案适合做教育类工具、作业批改系统、智能答题板的开发者,也适合想入门文档图像分析的算法同学。下面按数据、骨干、训练、部署、避坑的顺序讲透。
2. 数据从哪来:手写公式数据集的构建与预处理
2.1 为什么不能直接用印刷体公式数据集
印刷体公式数据集(比如用 LaTeX 渲染出来的图)和手写公式的分布差异极大。印刷体的笔画粗细均匀、字符间距固定、没有连笔和涂抹;手写体里同一个数字7可能带横杠也可能不带,+可能写成两笔交叉也可能一笔画完。如果拿印刷体数据训练再直接推理手写图,准确率会断崖式下跌。常见做法是:以公开手写数学符号数据集为底,叠加自己采集的笔迹数据做微调。采集时要注意覆盖不同书写习惯——左撇子、连笔、倾斜、不同笔宽。
2.2 图像预处理的四个关键步骤
预处理的目标是把任意尺寸、任意背景的笔迹图变成网络能吃的规整输入。下面这段代码是完整的预处理流水线,每一步都有明确的参数理由。
import cv2 import numpy as np def preprocess_formula_image(img_path, target_size=(224, 224)): # 1. 灰度化:手写公式不需要颜色信息,灰度能降通道、减计算量 img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) if img is None: raise ValueError(f"无法读取图像: {img_path}") # 2. 二值化:OTSU 自动找阈值,比固定阈值 127 更适应不同光照 # 反转使笔画为白(255)、背景为黑(0),符合后续归一化习惯 _, binary = cv2.threshold(img, 0, 255, cv2.THRESH_BINARY_INV + cv2.THRESH_OTSU) # 3. 去噪:中值滤波核 3x3,去掉孤立噪点但保留笔画边缘 denoised = cv2.medianBlur(binary, 3) # 4. 裁剪有效区域:找到所有非零像素的包围盒,去掉多余白边 coords = cv2.findNonZero(denoised) if coords is None: raise ValueError("图像中未检测到有效笔画") x, y, w, h = cv2.boundingRect(coords) cropped = denoised[y:y+h, x:x+w] # 5. 等比缩放 + 填充:保持宽高比,短边补零到目标尺寸 # 直接 resize 会拉伸字符,导致 '1' 变胖、'0' 变扁 h_orig, w_orig = cropped.shape scale = min(target_size[0] / w_orig, target_size[1] / h_orig) new_w, new_h = int(w_orig * scale), int(h_orig * scale) resized = cv2.resize(cropped, (new_w, new_h), interpolation=cv2.INTER_AREA) canvas = np.zeros(target_size, dtype=np.uint8) pad_x = (target_size[0] - new_w) // 2 pad_y = (target_size[1] - new_h) // 2 canvas[pad_y:pad_y+new_h, pad_x:pad_x+new_w] = resized # 6. 归一化到 [0,1] 并扩展通道维度,匹配 ResNet 输入格式 normalized = canvas.astype(np.float32) / 255.0 return np.expand_dims(normalized, axis=-1) # shape: (224,224,1)逻辑说明:二值化用 OTSU 而不是固定阈值,是因为手机拍照或平板截图的光照条件不可控,固定阈值在暗光下会把笔画和背景一起判黑。等比缩放加填充是手写识别里最容易被忽略的一步——很多人直接cv2.resize到 224×224,结果细长的公式被横向拉伸,1和l的区分度直接消失。归一化到 [0,1] 而不是 [-1,1],是因为后面接的是标准 ResNet 骨干,它的第一层卷积默认按 [0,1] 输入设计。
2.3 标签编码:从 LaTeX 到字符级序列
标签侧要做的是把\frac{1}{2}这种 LaTeX 串转成模型能学的序列。常见做法是维护一个字符表,覆盖 0-9、+、-、×、÷、(、)、=、x、y 以及分数线和根号等结构符号。编码时按字符切分,每个字符映射到一个整数 ID,序列两端加<sos>和<eos>。如果公式里出现字符表外的符号,直接丢弃该样本,不要用<unk>硬编码——手写公式里未知符号往往意味着标注错误,留着会污染训练集。
提示:字符表大小建议控制在 60 以内。超过这个数说明你在试图覆盖过于复杂的公式,而手写场景下复杂公式的标注一致性很难保证,不如先聚焦加减乘除和括号。
3. ResNet 骨干怎么选:从 ResNet18 到 ResNet50 的取舍
3.1 为什么是 ResNet 而不是 VGG 或 MobileNet
手写公式识别的骨干网络需要同时满足两个条件:足够深以捕捉二维结构关系,又不能太深导致小数据集过拟合。VGG 系列参数量大、没有残差连接,在几千张手写样本上训练很容易梯度消失;MobileNet 虽然轻,但深度可分离卷积对细笔画的特征提取偏弱,+和×这种细线交叉的区分度会下降。ResNet 的残差连接让梯度能直接回传,配合 BatchNorm 在小数据集上也能稳定收敛。实际选型时,ResNet18 是起步首选,ResNet34 是精度和速度的平衡点,ResNet50 只在样本量超过五万张时才有明显收益。
3.2 骨干网络的改造:输入通道与输出层
标准 ResNet 是为 RGB 三通道 224×224 设计的,手写公式是单通道灰度图,需要改第一层卷积。同时要去掉最后的 1000 类全连接层,换成适合序列输出的结构。
import torch import torch.nn as nn from torchvision import models class FormulaResNet(nn.Module): def __init__(self, num_classes, backbone='resnet18', pretrained=True): super().__init__() # 加载预训练骨干,pretrained=True 能加速收敛 if backbone == 'resnet18': self.backbone = models.resnet18(pretrained=pretrained) feat_dim = 512 elif backbone == 'resnet34': self.backbone = models.resnet34(pretrained=pretrained) feat_dim = 512 else: self.backbone = models.resnet50(pretrained=pretrained) feat_dim = 2048 # 改造第一层:3通道 -> 1通道 # 做法是把预训练权重的三通道求和,保留原始特征响应 old_conv = self.backbone.conv1 new_conv = nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3, bias=False) with torch.no_grad(): new_conv.weight.copy_(old_conv.weight.sum(dim=1, keepdim=True)) self.backbone.conv1 = new_conv # 去掉原始 fc 层,保留全局池化前的特征 self.backbone.fc = nn.Identity() # 接一个分类头做字符级预测(简化版,实际可用 CTC 或 Transformer 解码) self.classifier = nn.Sequential( nn.Linear(feat_dim, 256), nn.ReLU(inplace=True), nn.Dropout(0.3), nn.Linear(256, num_classes) ) def forward(self, x): # x: (B, 1, 224, 224) feat = self.backbone(x) # (B, feat_dim) logits = self.classifier(feat) # (B, num_classes) return logits逻辑说明:第一层卷积的权重改造是关键。直接把三通道权重随机初始化会浪费预训练特征,把三通道求和压成单通道,相当于把 RGB 三个通道的响应叠加到灰度上,保留了边缘和纹理检测能力。self.backbone.fc = nn.Identity()是 PyTorch 里去掉全连接层的标准写法,比手动改fc为None更安全,因为forward里还会调用它。分类头里的 Dropout 设 0.3 而不是 0.5,是因为手写公式的特征维度不算高,过强的 Dropout 会让训练震荡。
3.3 序列解码:CTC 还是 Attention
如果公式是单字符分类(比如只识别一个数字),上面的分类头就够了。但真实公式是变长序列,需要解码器。CTC(Connectionist Temporal Classification)适合字符间距不均匀、不需要显式对齐的场景,实现简单,但要求字符之间大致有序。Attention 解码器(Transformer Decoder)能处理二维结构,比如分数线的上下关系,但需要更多数据和更长的训练时间。我的经验是:先上 CTC 跑通 baseline,确认骨干特征有效,再换 Attention 做结构建模。不要一上来就堆 Transformer,否则调参周期会拖到让人怀疑人生。
注意:CTC 的 blank 标签要和字符表分开编号,通常设
num_classes = len(charset) + 1,最后一位是 blank。如果忘了加 blank,CTC 损失会直接报维度不匹配。
4. 训练与调参:让模型在几千张样本上收敛
4.1 数据增强:手写场景该做和不该做的
数据增强是手写识别里性价比最高的操作。该做的:随机旋转 ±10 度(模拟书写倾斜)、随机缩放 0.9~1.1(模拟不同字号)、弹性形变(模拟笔画抖动)。不该做的:水平翻转(会把+变成+但把2变成镜像,语义错误)、垂直翻转(6和9直接混淆)、大角度旋转(超过 15 度后公式结构被破坏)。下面是一个基于 albumentations 的增强配置。
import albumentations as A from albumentations.pytorch import ToTensorV2 train_transform = A.Compose([ A.Rotate(limit=10, border_mode=0, p=0.5), # 小角度旋转 A.RandomScale(scale_limit=0.1, p=0.3), # 缩放 0.9~1.1 A.ElasticTransform(alpha=1, sigma=50, p=0.2), # 弹性形变 A.GaussNoise(var_limit=(10, 50), p=0.2), # 模拟拍摄噪点 A.Normalize(mean=[0.5], std=[0.5]), # 单通道归一化 ToTensorV2() ])参数说明:Rotate的border_mode=0表示旋转后空白区域填 0(黑底),和预处理后的背景一致。ElasticTransform的alpha控制形变强度,1 是温和形变,超过 3 会让字符结构扭曲到不可读。GaussNoise的方差上限 50 是经验值,再高会淹没细笔画。
4.2 学习率与优化器的实际配置
ResNet 微调的学习率不能照搬 ImageNet 训练的 0.1。预训练权重已经包含通用特征,手写数据量小,学习率要降两个数量级。我一般用 AdamW,初始学习率 1e-4,权重衰减 1e-4,配合余弦退火。Batch size 在显存允许下尽量大,32 是底线,64 更稳。如果显存不够,用梯度累积模拟大 batch。
from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingLR model = FormulaResNet(num_classes=60, backbone='resnet18') optimizer = AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4) scheduler = CosineAnnealingLR(optimizer, T_max=50, eta_min=1e-6) # 训练循环核心片段 for epoch in range(50): model.train() for images, labels in train_loader: optimizer.zero_grad() logits = model(images) loss = nn.CrossEntropyLoss()(logits, labels) loss.backward() # 梯度裁剪防止梯度爆炸,阈值 5.0 是经验值 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0) optimizer.step() scheduler.step()逻辑说明:CosineAnnealingLR的T_max设成总 epoch 数,让学习率从 1e-4 平滑降到 1e-6。梯度裁剪的max_norm=5.0是手写识别里常用的保守值,因为弹性形变增强后的样本偶尔会产生异常梯度。如果训练 loss 在前 5 个 epoch 不下降,先检查数据标签是否对齐,再检查第一层卷积权重是否被正确初始化。
4.3 验证指标:准确率之外要看什么
字符级准确率会掩盖结构错误。比如1/2识别成12,字符准确率是 100%,但语义完全错了。验证时要同时看三个指标:字符错误率(CER)、公式级完全匹配率(EMR)、以及括号配对正确率。CER 用编辑距离算,EMR 要求整串完全一致,括号配对单独统计是因为括号错误在教育场景里最致命——(3+5少一个右括号,计算引擎直接报错。建议在验证集上每 5 个 epoch 打印一次这三个指标,不要只看 loss。
5. 避坑与排查:手写公式识别里最容易翻车的五件事
5.1 现象:训练 loss 正常下降但验证准确率卡在 30% 不动
原因:数据泄露或标签错位。常见情况是训练集和验证集里出现了同一张图的增强版本,模型在验证集上其实见过这些样本,但标签编码时字符表顺序不一致,导致预测 ID 和真实 ID 对不上。解决:先检查训练集和验证集的图像文件名是否有重叠,再打印前 10 个样本的标签解码结果,肉眼确认3没有被编码成8。
5.2 现象:模型对+和×的区分度极低,混淆矩阵里两者互相误判
原因:预处理阶段二值化阈值过高,细笔画被腐蚀。+和×都是细线交叉,如果 OTSU 把浅色笔画判成背景,两者的像素差异进一步缩小。解决:把二值化后的图存下来肉眼检查,如果笔画有断裂,改用自适应阈值cv2.adaptiveThreshold,或者把中值滤波核从 3 降到 1(相当于不做滤波)。
5.3 现象:推理时单张图耗时超过 500ms,批量推理又爆显存
原因:输入尺寸设成了 448×448 或更大,ResNet 的计算量随边长平方增长。手写公式的有效信息集中在笔画区域,224×224 足够覆盖大多数公式。解决:把输入尺寸降到 224×224,如果公式特别长(超过 15 个字符),改用 224×320 的矩形输入,而不是正方形放大。批量推理时用torch.no_grad()包住前向,显存占用能降一半。
5.4 现象:模型在测试集上表现好,但部署到实际场景后准确率暴跌
原因:训练数据的采集条件和实际使用条件不一致。训练集可能是扫描仪采集的白底黑字,实际场景是平板上的灰底蓝字,或者手机拍照有阴影。解决:在预处理里加一步自适应直方图均衡化(CLAHE),把不同光照条件下的对比度拉齐。同时收集至少 200 张实际场景的样本做微调,不要指望模型零样本泛化到新设备。
5.5 现象:括号识别正确但分数结构完全错乱,1/2被识别成12
原因:纯序列模型丢失了二维空间信息。CTC 或单向 Attention 只按从左到右的顺序解码,分数线的上下关系被压平。解决:在骨干特征后接一个二维位置编码,或者改用基于检测的方案——先检测分数线、括号等结构符号的位置,再按空间关系组装表达式树。如果不想改架构,至少在训练数据里加入大量分数样本,让模型从像素分布里隐式学到横线位置的含义。
6. 从识别到计算:表达式解析与结果验证的最后一公里
识别出 LaTeX 串只是中间产物,教育场景真正要的是计算结果。这一步的常见做法是:把 LaTeX 串解析成表达式树,再递归求值。解析时要注意运算符优先级和括号配对,求值时要把×和÷映射成*和/,同时处理除零和非法字符。
import re def latex_to_expr(latex_str): # 简化版:处理加减乘除和括号,分数用 '/' 表示 s = latex_str.replace('\\times', '*').replace('\\div', '/') s = s.replace('\\frac{', '(').replace('}{', ')/(').replace('}', ')') s = s.replace(' ', '') # 校验:只允许数字、运算符、括号和小数点 if not re.fullmatch(r'[0-9+\-*/().]+', s): raise ValueError(f"表达式含非法字符: {s}") # 括号配对检查 depth = 0 for ch in s: if ch == '(': depth += 1 elif ch == ')': depth -= 1 if depth < 0: raise ValueError("括号不匹配:右括号多余") if depth != 0: raise ValueError("括号不匹配:左括号多余") return s def safe_eval(expr): # 用 eval 前先做字符白名单校验,避免代码注入 allowed = set('0123456789+-*/().') if not set(expr).issubset(allowed): raise ValueError("表达式含不允许的字符") try: result = eval(expr, {"__builtins__": {}}, {}) except ZeroDivisionError: return "除数不能为零" return result # 示例 latex = r"\frac{1}{2} + 3 \times (4 - 1)" expr = latex_to_expr(latex) # 得到 "(1)/(2)+3*(4-1)" print(safe_eval(expr)) # 输出 9.5逻辑说明:latex_to_expr里的替换顺序很重要——先替换\times和\div,再处理\frac,否则\frac里的{}会干扰后续替换。括号配对检查用深度计数,遇到右括号时深度减到负数说明右括号多了,遍历结束深度不为零说明左括号多了。safe_eval里把__builtins__设为空字典,防止eval执行任意代码,这是处理用户输入时的基本安全习惯。
验证环节建议加一层反向校验:把计算结果代回原公式,检查是否满足等式。比如识别出3+5=8,计算左边得 8,和右边一致,说明识别和计算都正确。如果左边算出 9,要么是识别错了,要么是原题写错了,两种情况都值得记录日志,用于后续迭代模型。
我自己的习惯是:每次模型更新后,拿 50 张手写样本跑一遍端到端流程,从图像输入到计算结果输出,人工核对每一步。这比只看验证集准确率更能发现真实问题——验证集上的高分有时候只是过拟合的假象,端到端跑通才是硬道理。希望帮到你。
本文还有配套的精品资源,点击获取