1. 项目缘起从“炼丹”到“造轮子”的必经之路最近在复现手写数学公式识别领域的一篇经典论文——CANCompositional Attention Network时我遇到了一个几乎所有研究者都会面临的困境论文思路清晰开源代码也能跑通但一旦想换成自己的数据集就仿佛进入了一个布满暗礁的未知海域。官方代码往往是为特定数据集如 CROHME量身定制的从数据格式、预处理到模型输出每一个环节都可能藏着“坑”。网络上关于 CAN 的讨论要么停留在论文解读层面要么就是直接跑官方 Demo真正深入代码肌理、讲清楚如何适配私有数据集的实战分享少之又少。这促使我决定做一次彻底的代码梳理和改造。我的目标很明确不仅要理解 CAN 模型每一行代码在做什么更要打通从原始图片到最终 LaTeX 序列的完整训练流水线让它能“吃”进我自己的手写公式数据。这个过程远不止是改几个文件路径那么简单它涉及到数据加载器的重构、词表的管理、损失函数的适配甚至是解码策略的微调。如果你也正打算将 CAN 或其他类似的序列生成模型应用到自己的视觉文本识别任务上比如手写汉字、化学方程式或是任何需要从图像到序列映射的场景那么我踩过的这些坑、梳理出的这条路径或许能为你省下大量摸索的时间。2. 庖丁解牛深入 CAN 模型的核心架构与代码实现CAN 模型之所以在手写数学公式识别任务上表现出色核心在于其“组合式注意力”机制。它没有采用当时主流的 Encoder-Decoder 加全局注意力如 Bahdanau Attention的模式而是设计了一种更符合公式二维空间结构的注意力方式。简单来说传统注意力在解码每个符号时会去看编码器所有特征点的加权和而 CAN 的注意力是“分区域”的它试图先定位当前要识别的符号可能属于哪个大部件比如分式的分子、根号下的内容再在这个部件内部进行细粒度的特征聚焦。这种由粗到细的定位思想非常契合数学公式的递归树状结构。在代码层面我们通常拿到的开源实现例如基于 PyTorch 的版本会包含以下几个核心模块2.1 特征编码器Encoder编码器通常是一个深度卷积神经网络如 DenseNet 或 ResNet 的变种负责将输入的单通道灰度公式图像转换为一个二维的特征图。这个特征图的空间维度H, W相比原图大大缩小但通道数C很高蕴含了丰富的视觉语义信息。关键点在于编码器的输出需要保留足够的空间信息因为后续的注意力机制需要在这些空间位置上进行“巡览”。# 伪代码示意编码器结构 class Encoder(nn.Module): def __init__(self): super().__init__() # 例如使用一个轻量化的CNN backbone self.cnn nn.Sequential( nn.Conv2d(1, 64, kernel_size3, padding1), nn.BatchNorm2d(64), nn.ReLU(inplaceTrue), nn.MaxPool2d(2), # ... 更多卷积和池化层 ) def forward(self, x): features self.cnn(x) # 输出形状: [Batch, C, H, W] return features2.2 组合式注意力解码器Decoder这是 CAN 的灵魂。解码器是一个基于 RNN如 LSTM 或 GRU的序列生成模型但在每一步它执行两次注意力计算第一次注意力Group Attention预测一个“组注意力”权重用于在特征图上粗略地选择一个区域组。这可以理解为在决定当前要解码的符号属于公式的哪个逻辑部分。第二次注意力Within-group Attention在第一步选定的组内再进行一次细粒度的注意力计算得到最终用于解码的上下文向量。代码中你会看到两个独立的注意力模块它们的输入都是编码器特征图但权重计算方式不同。解码器 LSTM 的输入是上一步预测的字符嵌入或开始的sos标记与第二步得到的细粒度上下文向量的拼接。2.3 预测与训练解码器每一步会输出一个概率分布覆盖所有可能的符号包括数字、运算符、希腊字母、结构符号如\frac、\sqrt以及特殊的eos和pad。训练时我们使用交叉熵损失函数并以“教师强制”的方式将上一时刻的真实标签或前一步的预测结果输入给解码器。一个容易被忽略的细节是损失掩码。由于批处理时序列长度不一我们需要用pad填充到统一长度。计算损失时必须屏蔽掉这些填充位置否则会干扰模型学习。# 伪代码示意训练步骤的核心循环 for step in range(max_len): # teacher_forcing 决定使用真实标签还是自身预测 decoder_input ground_truth[:, step-1] if teacher_forcing else predictions.argmax(-1) group_attn_weights compute_group_attention(decoder_hidden, encoder_features) within_group_attn_weights compute_within_group_attention(decoder_hidden, encoder_features, group_attn_weights) context (within_group_attn_weights.unsqueeze(2) * encoder_features).sum(dim[1,2]) decoder_output, decoder_hidden decoder_lstm(torch.cat([char_embedding(decoder_input), context], dim1), decoder_hidden) step_prediction output_layer(decoder_output) loss cross_entropy_loss(step_prediction, ground_truth[:, step], ignore_indexpad_idx)理解了这个数据流你就掌握了 CAN 模型的运行脉搏。接下来我们要解决最实际的问题如何让这套代码为我们自己的数据集工作。3. 数据炼金术构建适配 CAN 的自定义数据集管道官方数据集如 CROHME通常提供预处理好的.pkl或特定格式的文本文件包含图像路径和对应的 LaTeX 序列。我们的私有数据往往是零散的图片和标注。构建数据管道是关键的第一步也是最容易出错的地方。3.1 数据准备与标注格式假设你有一批手写公式图片如formula_001.png,formula_002.jpg和一个记录了对应 LaTeX 公式的文本文件或 CSV。第一步是统一格式。我强烈建议创建一个dataset.csv包含两列image_path和latex_code。image_path,latex_code ./data/images/001.png, x \frac{-b \pm \sqrt{b^2 - 4ac}}{2a} ./data/images/002.jpg, \sum_{i1}^{n} i \frac{n(n1)}{2}注意LaTeX 代码中的空格需要特别注意。有些标注习惯在运算符周围加空格如\pm有些则不加。最好在预处理阶段进行统一比如移除所有非必要的空格或者制定一个明确的空格规则例如只在特殊符号和普通字符间加空格并在整个流程中保持一致否则词表会因同一个符号的不同表示而膨胀。3.2 实现自定义 Dataset 类这是 PyTorch 数据加载的核心。我们需要继承torch.utils.data.Dataset并实现三个方法__init__,__len__,__getitem__。import torch from torch.utils.data import Dataset from PIL import Image import pandas as pd import torchvision.transforms as T class HandwrittenFormulaDataset(Dataset): def __init__(self, csv_path, img_dir, transformNone, max_label_len150): Args: csv_path: 包含 image_path 和 latex_code 的 CSV 文件路径。 img_dir: 图片所在的根目录。 transform: 应用于图像的变换如缩放、归一化。 max_label_len: LaTeX 序列的最大长度用于填充。 self.df pd.read_csv(csv_path) self.img_dir img_dir self.transform transform self.max_label_len max_label_len # 词表将在外部构建并传入这里先留空 self.vocab None self.inv_vocab None def __len__(self): return len(self.df) def __getitem__(self, idx): img_name self.df.iloc[idx][image_path] latex self.df.iloc[idx][latex_code] # 1. 加载并处理图像 img_path os.path.join(self.img_dir, img_name) image Image.open(img_path).convert(L) # 转为灰度图 if self.transform: image self.transform(image) # 2. 将 LaTeX 字符串转换为索引序列 # 假设我们有一个将字符/单词映射到索引的词表 (self.vocab) # 以及特殊标记sos0, eos1, pad2, unk3 tokens self._tokenize_latex(latex) # 自定义分词函数 token_indices [self.vocab.get(token, self.vocab[unk]) for token in tokens] # 添加开始和结束标记 token_indices [self.vocab[sos]] token_indices [self.vocab[eos]] # 填充到统一长度 padded_indices token_indices [self.vocab[pad]] * (self.max_label_len - len(token_indices)) label torch.LongTensor(padded_indices[:self.max_label_len]) # 确保长度 # 3. 创建掩码可选用于损失计算 mask torch.zeros(self.max_label_len, dtypetorch.bool) mask[:len(token_indices)] 1 # 有效标记位置为1 return image, label, mask def _tokenize_latex(self, latex_str): 一个简单的 LaTeX 分词器示例。 更复杂的实现需要处理如 \frac{}{} 这样的命令。 # 这里使用一个简单策略按空格分割并分离反斜杠命令 # 例如\frac{a}{b} - [\frac, {, a, }, {, b, }] # 实际项目中可能需要更精细的语法解析器 tokens [] i 0 while i len(latex_str): if latex_str[i] \\: # 捕获命令直到第一个非字母字符 j i 1 while j len(latex_str) and latex_str[j].isalpha(): j 1 tokens.append(latex_str[i:j]) i j else: # 单独处理每个非命令字符包括花括号、空格等 if latex_str[i] ! or self.keep_space: # 根据策略决定是否保留空格 tokens.append(latex_str[i]) i 1 return tokens3.3 图像预处理与增强手写公式图像的大小、笔迹粗细、对比度差异很大。统一的预处理至关重要。缩放将图像缩放到固定高度如 64 像素宽度按比例缩放这是为了保持宽高比避免严重变形。然后可以将图像填充到固定宽度或者使用 CNN 编码器处理可变宽度。归一化将像素值从 [0, 255] 归一化到 [0, 1] 或 [-1, 1]。数据增强对于手写数据适度的增强可以提高模型鲁棒性。例如随机弹性形变模拟手写时的轻微抖动。轻微旋转±5度以内模拟扫描或拍摄时的角度偏差。对比度/亮度调整适应不同的纸张和墨水。添加椒盐噪声或高斯噪声模拟图像采集过程中的噪声。# 使用 torchvision.transforms 构建预处理流水线 from torchvision import transforms train_transform transforms.Compose([ transforms.Resize((64, 256)), # 固定高度设定一个较大的宽度用于填充 transforms.ToTensor(), transforms.Normalize(mean[0.5], std[0.5]), # 归一化到[-1,1] # 可以添加自定义的增强变换如 RandomElasticDistortion ]) # 验证集通常不需要增强 val_transform transforms.Compose([ transforms.Resize((64, 256)), transforms.ToTensor(), transforms.Normalize(mean[0.5], std[0.5]), ])3.4 构建词表Vocabulary词表是连接视觉特征和符号序列的桥梁。你需要遍历整个训练集的 LaTeX 标注收集所有独特的 token可能是字符级也可能是子词级或命令级。然后为每个 token 分配一个唯一的索引。def build_vocab(dataset, tokenizer_func, special_tokens[pad, sos, eos, unk]): 构建词表 vocab {} # 1. 添加特殊标记 for idx, token in enumerate(special_tokens): vocab[token] idx # 2. 收集所有 token all_tokens set() for latex in dataset.df[latex_code]: tokens tokenizer_func(latex) all_tokens.update(tokens) # 3. 为普通 token 分配索引 for token in sorted(all_tokens): # 排序保证可复现性 if token not in vocab: vocab[token] len(vocab) # 4. 创建反向词表 inv_vocab {v: k for k, v in vocab.items()} return vocab, inv_vocab # 使用示例 train_dataset HandwrittenFormulaDataset(train.csv, ./images) vocab, inv_vocab build_vocab(train_dataset, train_dataset._tokenize_latex) train_dataset.vocab vocab train_dataset.inv_vocab inv_vocab print(f词表大小: {len(vocab)})实操心得词表大小直接影响模型参数量和训练速度。对于数学公式命令如\frac,\sqrt是有限的但变量名如x,y,\alpha和数字组合可能很多。如果词表过大例如超过5000可以考虑采用子词切分如 BPE来压缩词表这对处理罕见或未登录的变量名特别有效。在初次尝试时可以先使用字符级词表每个字符作为一个 token虽然序列会变长但词表小简单可靠。4. 训练引擎改造让 CAN 模型“认识”你的数据有了数据管道下一步就是将数据喂给模型并调整训练脚本。官方代码的训练循环可能隐藏了许多针对原始数据集的假设我们需要逐一破解。4.1 修改数据加载部分找到训练脚本通常是train.py中加载数据的地方。将原来加载 CROHME 数据的代码替换为加载我们自定义的HandwrittenFormulaDataset。# 原代码可能类似这样假设 # from crohme_dataset import CROHMEDataset # train_set CROHMEDataset(...) # 替换为 from custom_dataset import HandwrittenFormulaDataset train_dataset HandwrittenFormulaDataset(csv_path./data/train.csv, img_dir./data/images, transformtrain_transform, max_label_len150) train_dataset.vocab vocab # 传入构建好的词表 val_dataset HandwrittenFormulaDataset(csv_path./data/val.csv, ...) val_dataset.vocab vocab # 创建 DataLoader train_loader DataLoader(train_dataset, batch_size32, shuffleTrue, num_workers4, collate_fncollate_fn) val_loader DataLoader(val_dataset, batch_size16, shuffleFalse, num_workers4, collate_fncollate_fn)4.2 实现 Collate 函数由于我们的图像高度固定、宽度可变如果采用保持宽高比的缩放一个批次内的图像宽度可能不同无法直接堆叠成张量。我们需要一个自定义的collate_fn来动态地将批次内所有图像填充到该批次的最大宽度。def collate_fn(batch): 处理可变宽度图像的批次组装 images, labels, masks zip(*batch) # 图像list of tensors [C, H, W_i] max_width max([img.shape[2] for img in images]) batch_size len(images) c, h images[0].shape[0], images[0].shape[1] padded_images torch.zeros(batch_size, c, h, max_width) for i, img in enumerate(images): padded_images[i, :, :, :img.shape[2]] img # 标签和掩码已经是填充好的可以直接 stack labels torch.stack(labels, dim0) masks torch.stack(masks, dim0) return padded_images, labels, masks4.3 调整模型输出层CAN 解码器的最后一层是一个线性层其输出维度等于原始词表的大小。现在我们的词表变了所以必须修改这个线性层。# 在模型初始化代码中找到 decoder 的定义 class Decoder(nn.Module): def __init__(self, hidden_size, encoder_feat_size, vocab_size, ...): super().__init__() # ... 其他层 ... self.fc_out nn.Linear(hidden_size, vocab_size) # 这里的 vocab_size 必须是我们新词表的大小 # 在实例化模型时传入新的 vocab_size vocab_size len(vocab) model CANModel(encoder..., decoderDecoder(..., vocab_sizevocab_size, ...), ...).to(device)4.4 损失函数与评估指标损失函数交叉熵本身不需要修改但要注意ignore_index参数必须设置为词表中pad标记的索引。评估指标则需要使用我们自己的词表进行解码。损失计算criterion nn.CrossEntropyLoss(ignore_indexvocab[pad])评估指标常用的有准确率Exact Match、BLEU 或 EDEdit Distance/Levenshtein Distance。在验证时我们需要将模型输出的索引序列通过inv_vocab转换回 LaTeX 字符串再与真实标签比较。注意比较前需要去除sos,eos,pad等特殊标记。def indices_to_latex(indices, inv_vocab): 将索引序列转换为 LaTeX 字符串 tokens [] for idx in indices: token inv_vocab[idx] if token eos: break if token not in [sos, pad]: tokens.append(token) # 将 token 列表拼接成字符串这里需要根据你的分词策略反向处理 # 例如如果分词时保留了空格可能需要用 连接 latex_str .join(tokens) # 或 .join(tokens) return latex_str def calculate_edit_distance(pred_indices, true_indices, inv_vocab): 计算编辑距离序列级别 pred_str indices_to_latex(pred_indices, inv_vocab) true_str indices_to_latex(true_indices, inv_vocab) # 使用 python-Levenshtein 库或自定义实现 import Levenshtein return Levenshtein.distance(pred_str, true_str)4.5 超参数调整与训练策略学习率这是最重要的超参数。对于新数据集建议从一个较小的学习率开始如 1e-4并使用学习率调度器如ReduceLROnPlateau或CosineAnnealingLR进行动态调整。批次大小受 GPU 内存限制。图像尺寸和序列长度是主要影响因素。如果内存不足可以尝试减小图像尺寸、缩短max_label_len或使用梯度累积。教师强制比率在训练早期使用较高的教师强制比率如 0.9有助于稳定训练。随着训练进行可以逐渐降低该比率让模型更多依赖自己的预测以提高推理时的鲁棒性。梯度裁剪对于 RNN/LSTM 解码器梯度爆炸是个常见问题。在optimizer.step()之前加入torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm5)是个好习惯。5. 实战排坑训练过程中的典型问题与解决方案即使按照上述步骤精心准备训练过程也绝不会一帆风顺。下面是我在实战中遇到的一些典型问题及其排查思路。5.1 Loss 不下降或震荡剧烈这是最令人头疼的情况。检查数据与标签对齐这是首要怀疑对象。随机抽取几个批次将图像可视化并将对应的标签索引用inv_vocab解码回 LaTeX 显示出来确保图片和公式是对应的。一个常见的错误是在数据清洗或分词时不小心引入了错位。检查词表覆盖验证集或训练集中是否出现了大量unk未登录词这会导致模型无法学习这些符号。检查构建词表时是否使用了全部训练数据分词函数是否正确处理了所有情况。学习率过高/过低绘制 Loss 曲线。如果 Loss 一开始就 NaN 或变得极大可能是学习率太高。如果 Loss 几乎不变可能是学习率太低或模型初始化有问题。尝试使用1e-4,1e-5等不同量级。梯度问题在训练循环中打印梯度的范数。如果梯度范数非常大或为 0说明存在梯度爆炸或消失。确保使用了梯度裁剪并检查网络层初始化例如LSTM 的初始化方式。损失掩码确认ignore_index是否正确设置为pad的索引。计算损失时如果没有屏蔽填充位Loss 会看起来在下降但模型实际在学无意义的东西。5.2 验证集准确率远低于训练集过拟合数据增强这是对抗过拟合最有效的手段之一。确保你的训练数据增强是合理且充分的。正则化在模型中加入 Dropout 层特别是在 LSTM 层之间和全连接层之前。也可以尝试权重衰减L2正则化。简化模型如果数据量很小而模型特别是编码器过于复杂如很深的 ResNet很容易过拟合。考虑使用更轻量的编码器如浅层 CNN或对预训练编码器进行更强的冻结。早停持续监控验证集指标如编辑距离当其在多个 epoch 内不再提升时停止训练。5.3 模型输出无意义或重复符号检查解码循环在推理非教师强制模式下解码器的输入是它上一步的预测结果。确保你正确地使用了model.eval()模式并且在每一步将argmax得到的索引作为下一步的输入而不是 softmax 后的概率。曝光偏差这是序列生成模型的通病。训练时使用真实标签教师强制推理时使用自身预测二者分布不一致。可以尝试计划采样在训练中逐步从教师强制过渡到自主预测。波束搜索贪婪解码每一步选概率最大的容易陷入局部最优。实现一个波束搜索Beam Search可以显著提升结果连贯性。虽然会减慢推理速度但对于数学公式这种结构严谨的序列效果提升通常很明显。5.4 训练速度慢数据加载瓶颈使用torch.utils.data.DataLoader的num_workers参数进行多进程数据加载。监控 GPU 利用率如果很低可能是 CPU 的数据预处理成了瓶颈。混合精度训练使用torch.cuda.amp进行自动混合精度训练可以大幅减少 GPU 内存占用并加快计算速度。检查图像尺寸过大的输入图像会导致编码器计算量激增。尝试将输入高度从 64 降到 48 或 32看看精度是否在可接受范围内下降。6. 超越训练模型推理、部署与效果优化当模型训练完成后工作只完成了一半。如何用它进行实际识别并进一步提升效果6.1 实现推理脚本创建一个独立的inference.py或predict.py。它需要加载训练好的模型权重、词表并对单张或一批图片进行预测。def predict(image_path, model, transform, vocab, inv_vocab, max_len150, devicecuda): model.eval() with torch.no_grad(): # 1. 图像预处理 image Image.open(image_path).convert(L) image_tensor transform(image).unsqueeze(0).to(device) # [1, C, H, W] # 2. 编码 encoder_outputs model.encoder(image_tensor) # 3. 解码贪婪或波束搜索 decoded_indices greedy_decode(model.decoder, encoder_outputs, vocab, max_len, device) # 或 decoded_indices beam_search_decode(model.decoder, encoder_outputs, vocab, max_len, device, beam_width5) # 4. 转换为 LaTeX latex_str indices_to_latex(decoded_indices, inv_vocab) return latex_str def greedy_decode(decoder, encoder_outputs, vocab, max_len, device): 贪婪解码示例 batch_size encoder_outputs.size(0) decoder_input torch.tensor([[vocab[sos]]] * batch_size, devicedevice) # [B, 1] decoder_hidden None decoded_ids [] for _ in range(max_len): # 这里需要根据你的 CAN 解码器前向传播方式调用 # 假设有一个 decode_step 函数 output, decoder_hidden, attn_weights decoder.decode_step(decoder_input, decoder_hidden, encoder_outputs) # output: [B, vocab_size] predicted_id output.argmax(-1) # [B] decoded_ids.append(predicted_id.item()) if predicted_id.item() vocab[eos]: break decoder_input predicted_id.unsqueeze(1) # 下一步的输入 return decoded_ids6.2 实现波束搜索贪婪解码每次选择概率最大的 token而波束搜索在每一步保留 top-K波束宽度个候选序列最终选择整体概率最高的序列。这能有效避免局部最优对于生成\frac{}{}这类需要正确配对括号和括号的命令尤其重要。def beam_search_decode(decoder, encoder_outputs, vocab, max_len, device, beam_width5): 简单的波束搜索实现 start_token vocab[sos] eos_token vocab[eos] # 初始波束 (序列, 对数概率, 解码器隐状态) beams [([start_token], 0.0, None)] # 初始隐状态为None for step in range(max_len): all_candidates [] for seq, log_prob, hidden in beams: if seq[-1] eos_token: # 如果序列已结束直接加入候选不再扩展 all_candidates.append((seq, log_prob, hidden)) continue decoder_input torch.tensor([[seq[-1]]], devicedevice) # 调用解码步获取下一步的概率和新的隐状态 output, new_hidden, _ decoder.decode_step(decoder_input, hidden, encoder_outputs) # output: [1, vocab_size] topk_probs, topk_ids torch.topk(output.squeeze(0), beam_width) for i in range(beam_width): next_token topk_ids[i].item() next_log_prob log_prob torch.log(topk_probs[i]).item() # 对数空间相加 new_seq seq [next_token] all_candidates.append((new_seq, next_log_prob, new_hidden)) # 从所有候选者中选择概率最高的 beam_width 个 ordered sorted(all_candidates, keylambda x: x[1], reverseTrue) beams ordered[:beam_width] # 检查是否所有波束都已结束 if all([seq[-1] eos_token for seq, _, _ in beams]): break # 返回概率最高的序列去掉开始标记 best_seq beams[0][0] return best_seq[1:] # 去掉 sos6.3 后处理与纠错模型输出的 LaTeX 字符串可能包含一些小的语法错误如括号不匹配、缺少空格等。可以编写简单的后处理规则进行纠正括号匹配检查遍历输出字符串检查{和}是否数量相等。命令补全检查\frac,\sqrt等命令后是否跟了必要的参数。空格标准化在运算符和操作数之间添加或移除空格使其符合 LaTeX 编译习惯。更高级的做法是训练一个小的序列纠错模型或者利用 LaTeX 编译器的错误信息进行反馈式修正。6.4 模型轻量化与部署如果考虑移动端或边缘设备部署需要对模型进行优化知识蒸馏用训练好的大模型教师去指导一个更小模型学生的训练。量化使用 PyTorch 的量化工具将模型权重从 FP32 转换为 INT8大幅减少模型体积和推理延迟精度损失通常很小。ONNX 导出将模型导出为 ONNX 格式便于在不同推理引擎如 TensorRT, OpenVINO上部署。整个流程走下来从理解论文、梳理代码到构建数据管道、调试训练、优化推理是一个完整的深度学习项目闭环。其中最大的收获不是调出了一个高精度的模型而是掌握了将一篇学术论文的创意落地解决一个实际工程问题的全套方法论。当你看到自己手写的潦草公式被模型准确地转换成工整的 LaTeX 代码时那种成就感远非单纯跑通一个 Demo 可比。这个过程教会你的是如何让代码服务于你的需求而不是被代码所束缚。