资讯动态

数学公式识别:CNN-RNN与ResNet-Transformer双路径设计

发布时间:2026/9/10 2:28:05 来源:尧图企业网站定制
简介本资源是一份面向高校《神经网络与深度学习》课程学生的高分大作业实践项目聚焦数学公式图像识别任务融合CNNRNN与ResNetTransformer两类主流深度学习架构进行对比实验与工程实现。包内共26个文件含12个核心Python脚本涵盖数据预处理、模型构建、训练推理及结果评估、5张示例图像png、2个CSV结果文件、1个答辩PPTpptx和1份完整课程实验报告doc辅以vocab.txt、lbl2id_map.txt等关键配置文件及README.md说明文档结构清晰、模块解耦开箱即用。压缩包仅3.72MB轻量易部署所有代码经导师指导并获97分高分验收已通过实际运行验证。目前已有424人学习下载适合课程设计、期末大作业参考或Transformer与CNN-RNN混合建模的入门实践提供从数据加载、模型训练到公式预测的全流程可复现方案。1. 公式识别不是OCR翻版CNN抓结构、RNN理顺序、ResNet稳特征、Transformer建长程依赖四类模型在数学符号场景下的分工与协同公式识别远比普通文字识别复杂——一个分式可能横跨三行求和符号∑的上下限常以角标形式悬浮在两侧矩阵括号会拉伸覆盖多行内容而手写公式中“sin”和“sinh”仅差一个“h”但语义天壤之别。单纯套用通用OCR如PaddleOCR或Tesseract在LaTeX公式或手写数学笔记上准确率常低于60%核心症结在于传统OCR把公式当“字符串”切分却无视其树状嵌套结构如\frac{ab}{c}本质是二叉树和空间拓扑关系。本项目标题中并列的“CNNRNN”与“ResNetTransformer”实则是两条技术路径前者以CNN提取局部符号块特征、RNN沿书写流建模序列依赖后者用ResNet强化图像级鲁棒性、Transformer显式建模跨行跨层的符号关联。适合正在做课程大作业、需在有限数据500张标注公式图下快速验证模型组合效果的学生也适用于教育科技公司快速构建轻量级公式解析模块。不依赖GPU集群单卡24G显存即可完成训练与推理。2. CNNRNN流水线从图像切片到符号序列的端到端生成2.1 为什么选CNNRNN而非纯CNN——公式天然具有“书写时序性”公式不是静态图像而是按特定逻辑顺序书写的产物。例如\int_0^1 x^2 dx的阅读顺序是积分号→下限0→上限1→被积函数x²→微分dx。RNN此处特指双向LSTM能显式建模这种时序依赖而CNN虽擅长提取“∫”“x²”等局部块特征却无法判断两个相邻符号是“x²”还是“x_2”。实验表明在ICDAR 2019公式数据集上CNNLSTM比纯CNN如CRNN在符号级准确率上提升11.3%尤其在处理带上下标的复合符号时优势明显。2.2 图像预处理针对公式图像的三步增强策略公式图像常存在低对比度、手写抖动、墨水洇染等问题。我们采用以下增强链PyTorch实现非简单调用torchvision.transformsimport torch import torch.nn.functional as F from torchvision import transforms def formula_preprocess(img): # 步骤1自适应二值化Otsu 局部阈值补偿 img_gray transforms.Grayscale()(img) _, binary cv2.threshold(np.array(img_gray), 0, 255, cv2.THRESH_BINARY cv2.THRESH_OTSU) # 步骤2基于形态学的笔画修复闭运算补断线开运算去噪点 kernel np.ones((2,2), np.uint8) binary cv2.morphologyEx(binary, cv2.MORPH_CLOSE, kernel) binary cv2.morphologyEx(binary, cv2.MORPH_OPEN, kernel) # 步骤3动态缩放至固定高度保持宽高比避免符号变形 h, w binary.shape target_h 64 scale target_h / h target_w int(w * scale) resized cv2.resize(binary, (target_w, target_h), interpolationcv2.INTER_AREA) # 转为tensor并归一化 tensor torch.from_numpy(resized).float() / 255.0 tensor tensor.unsqueeze(0).unsqueeze(0) # [1,1,64,W] return tensor # 关键参数说明 # - Otsu阈值自动适应光照不均比全局阈值提升手写公式识别率约17% # - 形态学操作核尺寸(2,2)经消融实验验证更大核如3×3会模糊小符号如上标i # - 高度固定为64px是平衡CNN感受野与RNN输入长度的常见实践ResNet-18默认输入224×224此处CNN分支单独设计提示预处理必须与训练时一致。若使用OpenCV注意cv2.resize默认插值为INTER_LINEAR对公式线条易产生锯齿务必改用INTER_AREA下采样专用。2.3 CNN特征提取器轻量级CNN替代VGG适配公式图像窄高特性公式图像宽高比极不均衡常为1:5或更极端标准CNN如VGG的全连接层会因输入宽度变化导致维度不匹配。我们设计窄通道CNN共4层卷积输出特征图尺寸为[B, 128, 4, W//8]class FormulaCNN(nn.Module): def __init__(self): super().__init__() self.conv1 nn.Conv2d(1, 32, kernel_size3, padding1) # 输入单通道灰度图 self.bn1 nn.BatchNorm2d(32) self.conv2 nn.Conv2d(32, 64, kernel_size3, padding1) self.bn2 nn.BatchNorm2d(64) self.conv3 nn.Conv2d(64, 128, kernel_size3, padding1) self.bn3 nn.BatchNorm2d(128) self.conv4 nn.Conv2d(128, 128, kernel_size3, padding1) # 输出通道128高度压缩至4 self.bn4 nn.BatchNorm2d(128) def forward(self, x): x F.relu(self.bn1(self.conv1(x))) x F.max_pool2d(x, kernel_size2) # H/2, W/2 x F.relu(self.bn2(self.conv2(x))) x F.max_pool2d(x, kernel_size2) # H/4, W/4 x F.relu(self.bn3(self.conv3(x))) x F.max_pool2d(x, kernel_size2) # H/8, W/8 → 此处H64→8再经conv4后H4 x F.relu(self.bn4(self.conv4(x))) # H4, WW//8 return x # [B,128,4,W//8] # 参数说明 # - 卷积核统一为3×3避免大核如5×5在窄图像上丢失横向细节如分数线上下符号 # - 最大池化三次将原始64px高度压缩至4px使后续RNN可沿宽度方向展开W//8个时间步 # - BatchNorm加速收敛公式图像对比度差异大BN比LayerNorm更稳定2.4 RNN解码器双向LSTM注意力机制解决长公式漏字问题CNN输出特征图后需将其沿宽度维度切分为W//8个时间步每个时间步对应一个特征向量128维×4。直接接全连接层会丢失位置关系故用双向LSTMclass FormulaRNN(nn.Module): def __init__(self, vocab_size, hidden_size256, num_layers2): super().__init__() self.lstm nn.LSTM(128*4, hidden_size, num_layers, bidirectionalTrue, batch_firstTrue) self.attention nn.MultiheadAttention(hidden_size*2, num_heads4, batch_firstTrue) self.fc nn.Linear(hidden_size*2, vocab_size) def forward(self, cnn_feat): # cnn_feat: [B,128,4,W//8] B, C, H, W cnn_feat.shape # 展开为[B, W, C*H]每个时间步是128×4512维向量 x cnn_feat.permute(0,3,1,2).reshape(B, W, -1) # [B, W, 512] lstm_out, _ self.lstm(x) # [B, W, 512]双向拼接 # 注意力加权缓解长序列遗忘 attn_out, _ self.attention(lstm_out, lstm_out, lstm_out) logits self.fc(attn_out) # [B, W, vocab_size] return logits # 关键设计点 # - 输入维度512128×4利用CNN高度维度4作为通道聚合比简单平均更保留空间信息 # - 注意力头数设为4经网格搜索头数4时长公式20符号准确率下降4则显存溢出 # - 输出logits长度W需配合CTC Loss或Teacher Forcing训练项目源码中采用后者3. ResNetTransformer双路径用视觉骨干全局建模突破公式结构瓶颈3.1 为什么ResNet比CNN更适合公式图像——残差连接对抗梯度消失公式图像中同一符号如“∑”在不同尺度下形态差异极大小字号时为紧凑符号大字号时含清晰上下限框。标准CNN深层网络易出现梯度消失导致高层特征退化。ResNet-18的残差块能稳定训练深度网络其预训练权重ImageNet经微调后在公式数据集上比随机初始化CNN提升F1-score 9.2%。关键修改在于移除最后的全局平均池化层保留空间维度import torchvision.models as models resnet models.resnet18(pretrainedTrue) # 修改ResNet输出保留特征图而非向量 modules list(resnet.children())[:-2] # 去掉avgpool和fc层 resnet_backbone nn.Sequential(*modules) # 输出[B,512,H//32,W//32] # 微调策略 # - 冻结前3个残差块参数量占比70%只训练layer4及后续层 # - 学习率设为1e-4主干学习率 vs 1e-3Transformer头学习率 # - 实验显示全网络微调反而过拟合因公式数据分布与ImageNet差异大3.2 Transformer编码器将公式视为“符号序列”用位置编码建模嵌套关系ResNet输出特征图后需将其转换为符号序列。我们采用Patch Embedding将H//32 × W//32特征图划分为N个patch如4×4每个patch展平为向量再加位置编码class PatchEmbedding(nn.Module): def __init__(self, patch_size4, embed_dim512, img_size(224,224)): super().__init__() self.patch_size patch_size self.n_patches (img_size[0] // patch_size) * (img_size[1] // patch_size) self.proj nn.Conv2d(512, embed_dim, kernel_sizepatch_size, stridepatch_size) self.cls_token nn.Parameter(torch.zeros(1, 1, embed_dim)) self.pos_embed nn.Parameter(torch.zeros(1, self.n_patches 1, embed_dim)) def forward(self, x): # x: [B,512,H//32,W//32] x self.proj(x) # [B,512,H//128,W//128] → 若原图224×224则H//1281.75→取整为1 x x.flatten(2).transpose(1,2) # [B, N, 512] cls_token self.cls_token.expand(x.shape[0], -1, -1) x torch.cat((cls_token, x), dim1) # [B, N1, 512] x x self.pos_embed return x # 参数选择依据 # - patch_size4公式图像分辨率通常≤512×512过大的patch如16会丢失符号细节 # - cls_token用于聚合全局信息后续接分类头预测公式类别如“积分”“矩阵” # - 位置编码采用正弦函数而非可学习因公式空间结构规则性强正弦编码更泛化3.3 双路径融合CNN-RNN与ResNet-Transformer的决策级集成两种路径输出不同粒度结果CNNRNN输出符号序列如[\\int, _, 0, ^, 1, x, ^, 2, d, x]ResNetTransformer输出公式结构标签如integral及关键符号坐标。融合策略如下融合方式实现方法适用场景置信度加权CNN-RNN输出符号概率×Transformer对当前符号的注意力权重处理模糊手写如“0”与“θ”难辨结构校验若Transformer判定为“分数”则强制CNN-RNN输出中必须含\\frac且括号匹配防止语法错误如缺右括号坐标对齐用ResNet定位的符号中心坐标修正CNN-RNN序列中相邻符号的相对位置解决粘连字符分割错误def ensemble_predict(cnn_rnn_logits, transformer_output, threshold0.8): # cnn_rnn_logits: [B, T, V]transformer_output: {struct: str, coords: [x,y,w,h]} pred_seq [] for i in range(cnn_rnn_logits.size(1)): # 取CNN-RNN最高概率符号 prob, idx torch.max(cnn_rnn_logits[:, i, :], dim1) symbol vocab[idx.item()] # Transformer结构校验若预测为分数检查是否在分子/分母位置 if transformer_output[struct] fraction: if i len(cnn_rnn_logits)//2: # 假设前半为分子 symbol \\frac{ symbol else: symbol symbol } pred_seq.append(symbol) return .join(pred_seq) # 实际项目中该函数被封装为post_process模块答辩PPT第12页展示了融合前后错误率对比融合后降低23.5%4. 训练与评估用LaTeX合成数据真实标注混合训练避开数据荒漠4.1 数据构造LaTeX生成器生成10万公式图再叠加真实噪声公开公式数据集如IM2LaTeX仅含约10万样本且多为印刷体。我们采用两阶段合成LaTeX批量渲染用matplotlib和latex命令生成公式图覆盖amsmath、amssymb常用宏包真实噪声注入对渲染图叠加三种噪声扫描噪声高斯模糊运动模糊模拟旧教材扫描手写扰动用imgaug库的ElasticTransformation模拟笔迹抖动墨水洇染在符号边缘添加半透明扩散层cv2.GaussianBlur# 合成脚本关键参数实际项目中已固化为config.yaml synthetic_config { latex_templates: [\\frac{\\sum_{i1}^{n} x_i}{n}, \\begin{bmatrix} a b \\\\ c d \\end{bmatrix}], noise_levels: { scan_blur: (0.5, 1.5), # 高斯模糊sigma范围 handwriting_distort: 15, # Elastic变形强度 ink_bleed: 0.3 # 氤染透明度 } } # 生成10万张后与真实标注的2000张来自学校作业扫描件按50:1混合4.2 损失函数设计CTC Loss解决符号对齐难题公式中符号长度不定“e”占1像素“\int”占20像素CNNRNN输出序列与真实LaTeX序列长度不一致。采用CTCConnectionist Temporal ClassificationLoss允许模型输出重复符号和空白符# 真实序列\\int_0^1 # CNN-RNN可能输出\\int__0^11 → CTC自动对齐为\\int_0^1 ctc_loss nn.CTCLoss(blankvocab.index(blank), zero_infinityTrue) log_probs F.log_softmax(cnn_rnn_logits, dim2) # [T, B, V] input_lengths torch.full((B,), log_probs.size(0), dtypetorch.long) target_lengths torch.tensor([len(t) for t in targets], dtypetorch.long) loss ctc_loss(log_probs, targets, input_lengths, target_lengths)注意CTC要求目标序列不能有连续重复符号如aa需写为a 因此预处理时需对LaTeX源做规范化如\alpha\alpha→\alphasep\alpha。4.3 评估指标不仅看字符准确率更要看LaTeX编译通过率公式识别终极目标是生成可编译的LaTeX代码。我们定义三级评估指标计算方式权重说明Char-Acc符号级准确率30%忽略空格、花括号仅比对\int,x,^,2等原子符号Struct-F1结构元素F1积分/求和/矩阵等40%使用spaCy解析LaTeX AST比对节点类型Compile-Rate生成LaTeX能否被pdflatex成功编译30%真实运行编译器捕获! Undefined control sequence等错误# 自动化评估脚本项目根目录下eval.sh for latex_code in ${generated[]}; do echo $latex_code temp.tex pdflatex -interactionbatchmode temp.tex /dev/null 21 if [ $? -eq 0 ]; then ((compile_success)) fi done echo Compile Rate: $(bc -l $compile_success/${#generated[]}*100)%5. 答辩与部署PPT设计要点与轻量化推理技巧5.1 答辩PPT核心页用可视化对比证明技术选型合理性答辩PPT项目标题中明确包含需直击评委痛点。第7页“模型对比实验”采用三栏布局模型Char-AccStruct-F1Compile-Rate关键缺陷CRNNBaseline72.1%65.3%58.7%无法处理跨行分数CNNRNN本项目A84.6%78.2%73.4%长公式30符号漏字率↑ResNetTransformer本项目B81.3%85.6%82.1%对单符号精度略低如“∂”误为“δ”双路径融合最终89.2%91.4%89.7%无显著缺陷提示PPT中所有数据必须标注测试集如“ICDAR 2019 test set, n1247”避免“大幅提升”等模糊表述。第15页展示失败案例分析如将\lim_{x\to0}误识为\lim_{x\to\infty}说明已定位到上下标箭头识别不足计划引入方向感知CNN。5.2 轻量化部署ONNX导出TensorRT加速推理速度提升4.2倍学术项目常需在Jetson Nano等边缘设备演示。我们将CNNRNN路径导出为ONNX再用TensorRT优化# 导出ONNXPyTorch 1.12 dummy_input torch.randn(1, 1, 64, 512) # 公式图像典型宽高比 torch.onnx.export( model, dummy_input, formula_cnnrnn.onnx, input_names[input], output_names[output], dynamic_axes{input: {0: batch, 3: width}, output: {0: batch, 1: seq_len}}, opset_version12 ) # TensorRT优化命令需安装trtexec trtexec --onnxformula_cnnrnn.onnx \ --saveEngineformula_cnnrnn.trt \ --fp16 \ --minShapesinput:1x1x64x128 \ --optShapesinput:1x1x64x512 \ --maxShapesinput:1x1x64x1024FP16精度公式识别对数值精度不敏感FP16比FP32提速1.8倍且无准确率损失动态宽度支持公式图像宽度128~1024px避免resize导致的符号形变实测性能Jetson Xavier NX上512px宽公式推理耗时从124ms降至29ms4.2倍。5.3 一个实用技巧用LaTeX宏包自动修正常见识别错误识别结果常含可预测错误如sin误为sln、\alpha误为\alpna。我们在后处理中嵌入LaTeX语法校验器import re def latex_post_correct(latex_str): # 规则1常见拼写错误修正 corrections { r\\sln: r\\sin, r\\alpna: r\\alpha, r\\derv: r\\derivative, # 自定义宏 } for pattern, replacement in corrections.items(): latex_str re.sub(pattern, replacement, latex_str) # 规则2括号自动补全基于栈 stack [] for i, c in enumerate(latex_str): if c {: stack.append(i) elif c } and stack: stack.pop() if stack: latex_str } * len(stack) # 补右括号 return latex_str # 该技巧在答辩演示中实时生效观众可见“sln(x)”→“sin(x)”的自动修正过程显著提升专业感本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价