基于PyTorch的数学公式识别:卷积编码器与LSTM注意力解码器实战
2026/9/14 2:49:28 网站建设 项目流程

简介:基于神经网络模型的数学公式识别毕业设计项目,面向计算机视觉与深度学习方向的学生,可用来理解公式识别原理、完成课程设计或快速搭建评估环境。项目源码已本地编译可运行,评审分达95分以上,难度适中,整体共76个文件,压缩包约44.5MB。代码以35个Python脚本为主,覆盖模型构建、训练、预测与评估等环节;另有JSON配置/词表、TXT公式数据集、ipynb分析笔记及docx文档说明,便于对照复现。项目包含编码器-解码器结构、注意力可视化和序列生成等典型模块,目录按模型、工具、配置、评估等划分,结构清晰;动图演示可直观查看公式识别与注意力权重变化。已有145人学习下载,适合希望在数学公式识别方向上快速获得可运行参考实现并深入理解细节的中高级学习者。

1. 公式识别比普通 OCR 难在哪:先理解“为什么是序列问题”

先把结论放前面:数学公式识别不是“读出图里的字”,而是从二维版面结构里还原 LaTeX 语义。普通 OCR 按行切分再接字典,遇到根号、上下标、分式就错乱,因为关键信息藏在嵌套和相对位置里。近几年的主流做法是把公式图像送进卷积神经网络编码器,把 LaTeX 字符串作为序列解码器的目标,端到端训练一个“图到序列”模型。这个思路对毕业设计尤其合适:模型有标准版型、公开数据集可下载、PyTorch 代码全都可控。下面按任务建模、Python 实现、训练推理、交付组织四个部分讲清楚,源码和文档说明该长什么样,最后一章给具体模板。

2. 公式识别网络怎么搭:卷积编码器 + 自回归解码器

2.1 把公式识别重新定义为“图到序列”任务

公式图像和普通文本图像的本质区别:普通文本的语义顺序和视觉顺序一致,可以按行扫描;公式的语义顺序由结构决定,分式的分子和分母在图像上是上下排列,但在 LaTeX 串里却是\frac{分子}{分母}这样的线性顺序,视觉位置和输出顺序并不天然对应。

早期方案分三步:版面分析、字符切分、结构解析。每一步都是独立模型或规则,误差逐层累积,切分错一个符号,后面的结构解析就全乱。端到端模型把这三步合成一步:输入图像,输出 LaTeX 序列,中间的所有对齐关系由神经网络自己学,省掉了大量人工设计规则。这套范式从 2016 年前后的论文开始成为主流,到现在依然是绝大多数公式识别系统的基线架构。

任务可以形式化成一个条件概率:给定图像 x,求 LaTeX 序列 y 的条件概率,解码时逐 token 生成 y。目标序列的词表是 LaTeX 里出现的字符和控制序列,输出端就是一个自回归语言模型。把公式识别放进这个框架后,训练代码可以直接借用机器翻译的流程,区别只在编码器一侧。这也解释了为什么“数学公式识别 神经网络模型 Python”这几个关键词组合在一起时,搜到的资料大多是序列到序列的框架。

2.2 编码器:卷积神经网络负责压缩二维版面

编码器直接用 ImageNet 预训练的卷积神经网络模型,常见选择是 ResNet-18 或 ResNet-50。输入是一张 224×224 的灰度公式图,为了配合预训练权重,需要在输入端复制成 3 通道。经过 ResNet 的卷积和池化层,特征图逐步降到原图尺寸的 1/32,也就是 7×7 的空间分辨率。每个空间位置的 512 维特征向量再用 1×1 卷积投影到 d_model 维,最后把 49 个位置向量按顺序拼成一个序列。这个阶段相当于用卷积神经网络把“公式版面”压缩成一组空间位置上的语义向量,后面的解码器只能看到这些向量,不再直接接触原始像素。

为什么不用更深的网络?公式识别的难点在结构关系,不在纹理细节。ResNet 残差结构训练稳定,显存占用也友好,8G 显卡能带上 32 的 batch size。如果担心 7×7 分辨率丢细节,可以在 ResNet 最后一层之前截断,或者移除最后一个 stride=2 的下采样,把特征分辨率提高到 14×14,代价是显存增加,这属于后置调优手段。下面这段形状对照表就是编码器内部的张量流:

# 输入输出形状对照(B 为 batch size) # x : (B, 3, 224, 224) 灰度图复制为三通道 # feat : (B, 512, 7, 7) ResNet-18 特征 # feat : (B, 256, 7, 7) 1x1 卷积投影到 d_model # seq : (B, 49, 256) 展平成序列,N = 7*7

这段对照的价值在调参时很直接:改输入尺寸、改 ResNet 层数、改 d_model,先看形状是否符合预期,再去跑训练。1×1 卷积在这里有两重作用:统一特征维度,方便和后续解码器隐状态拼接;同时对每个位置的通道表达做一次重新加权。如果跳过这一层,解码时的注意力拼接维度会不匹配,代码上会频繁报维度错误。

2.3 解码器的自回归角色与注意力对齐

解码器有两个主流选择:LSTM 加注意力,或者 Transformer。毕业设计选型上我一般会推 LSTM 加注意力:结构直观,论文里好解释,逐 token 生成的过程和人类写 LaTeX 的思维一致,attention 可视化也顺手。Transformer 是编码解码全并行加自注意力,训练时更快,但需要更大的数据和更精细的学习率调度,数据量不到十万级时收敛不如 LSTM 稳。

自回归的意思是:生成第 t+1 个 token 时,会把已经生成的前 t 个 token 作为条件输入。训练阶段用 teacher forcing,直接把真实标签里的前 t 个 token 当输入,模型每步看到的都是正确历史;推理阶段只能拿自己每一步的预测结果继续往后生成,两种状态下历史分布不同,这就是训练和推理不一致的问题来源。缓解手段包括计划采样(teacher forcing 概率随 epoch 衰减)、标签平滑、以及束搜索。

注意力在公式识别里还有额外价值:解码器输出\frac时,注意力权重应该集中在图像的分式横线区域;输出\sqrt时,注意力应该聚到根号左上角和被开方区域。把注意力热力图叠加到原图上,可以直接判断模型是否真的“看”对了位置,这个可视化本身就是答辩时很有力的中间结果。

2.4 损失函数与评价指标的选择

训练用交叉熵,target 序列做 padding 时要把 padding 位置在损失里 ignore 掉,否则模型会学着输出大量[PAD]。PyTorch 的 CrossEntropyLoss 自带 ignore_index 参数,把 padding token id 传进去即可。

评价指标上,公式识别不太用普通 OCR 的整句准确率。字符级编辑距离最常用:它给“只错一个符号”的结果一个小惩罚,给“结构全错”的结果一个大惩罚,能量化模型离正确答案差多远。比编辑距离更严格的是公式级识别率,整条 LaTeX 完全一致才算对。这里有个细节:如果用字符级 tokenizer,\frac会被拆成五个字符,编辑距离会高估错误;实践里一般把 LaTeX 控制序列切成一个 token,再计算编辑距离或报告符号正确率。这几个指标在评估脚本里分别统计,文档里分开报,能明显看出模型是结构错误多还是符号错误多。

提示:CROHME 这类公开基准报告的主要指标就是公式级识别率,常规模型在这个数上到 50% 以上就可以写进文档,不必追求 90% 这种离谱数字。

3. 用 Python 从零搭起可训练的公式识别模型

3.1 数据集准备:CROHME 与合成样本两条路

公式识别数据集分两类。一类是公开的手写公式数据集,CROHME 最常用,包含在线笔迹的 InkML 文件和对应的 LaTeX 标注,科研论文的 benchmark 也是它。另一类是合成数据:写一批 LaTeX 公式模板,随机替换符号和数字,渲染成图片,可以无限扩量,是毕业设计里提升识别率的常见做法。两条路可以混合:先用公开数据练到基线,再用合成数据补生僻符号。

如果走合成路线,渲染这一步可以用 pdfLaTeX 加 ImageMagick 批量完成,命令大致如下:

# formulas/ 下每一个 .tex 文件是一条单行公式 for f in formulas/*.tex; do pdflatex -interaction=nonstopmode "$f" >/dev/null 2>&1 convert "${f%.tex}.pdf"[0] -trim -resize 224x224 -gravity center \ -background white -extent 224x224 "images/$(basename "${f%.tex}").png" done

这个脚本把每条 LaTeX 渲染成 PDF,再转成 224×224 的白底灰度 PNG。参数上-trim去掉公式周围白边,-extent统一画布尺寸,避免不同公式的空白差异影响模型;-resize按较短边等比缩放,短公式不会被拉伸变形。如果已经拿到了公开数据集的渲染图,这步直接跳过。数据量上,起步 5000 张能跑通流程,想看到明显的识别效果建议扩到 2 万张以上。

3.2 字符级 tokenizer 与数据加载器

LaTeX 标注需要转成 token id。最简单的做法是字符级切分,词表就是常见 ASCII 字符加\{},再加四个特殊 token。下面这个 tokenizer 类可以直接抄:

from collections import Counter class CharTokenizer: def __init__(self, corpus_lines, min_freq=1): self.char2idx = {'[PAD]': 0, '[SOS]': 1, '[EOS]': 2, '[UNK]': 3} counter = Counter() for line in corpus_lines: counter.update(line.strip()) for ch, freq in counter.items(): if freq >= min_freq: self.char2idx.setdefault(ch, len(self.char2idx)) self.idx2char = {v: k for k, v in self.char2idx.items()} def encode(self, text): ids = [self.char2idx['[SOS]']] ids += [self.char2idx.get(ch, self.char2idx['[UNK]']) for ch in text] ids.append(self.char2idx['[EOS]']) return ids def decode(self, ids): chars = [self.idx2char.get(i, '?') for i in ids] return ''.join(chars).replace('[EOS]', '').replace('[PAD]', '').replace('[SOS]', '')

参数说明:min_freq用来过滤整个语料里只出现一两次的冷僻字符,过滤掉的统一映射到[UNK],这样词表不会出现“只见过一次的 token”导致过拟合。个人经验是把 min_freq 从 1 调到 2,错误率基本不变,但词表能小 10%,训练更快。字符级缺点是把控制序列拆散,进阶做法是先把\frac\sqrt\sum这类整体映射成一个 token,词表多一两百个,解码错误会明显减少。

数据加载器沿用 PyTorch Dataset 接口,重点是保持“图片 + LaTeX 文本”成对返回:

from PIL import Image from torch.utils.data import Dataset class FormulaDataset(Dataset): def __init__(self, image_dir, label_file, tokenizer, transform=None): self.image_dir = image_dir self.tokenizer = tokenizer self.transform = transform self.samples = [] with open(label_file, 'r', encoding='utf-8') as f: for line in f: img_name, latex = line.rstrip().split('\t') self.samples.append((img_name, latex)) def __len__(self): return len(self.samples) def __getitem__(self, idx): img_name, latex = self.samples[idx] img = Image.open(f"{self.image_dir}/{img_name}").convert('RGB') if self.transform: img = self.transform(img) return img, self.tokenizer.encode(latex)

label_file每行格式是“图片文件名 + 制表符 + LaTeX 源码”。convert('RGB')是为了契合预训练卷积的 3 通道输入。返回值第二项长度不固定,到 DataLoader 汇聚成 batch 时需要 collate_fn 做 padding,同时记录每条真实长度,后面解码时按真实长度截断。image 端标准化用 ImageNet 的 mean/std,因为编码器加载的是预训练权重。

3.3 模型定义:编码器与解码器

编码器把 ResNet 特征投影成序列,代码实现如下:

import torch.nn as nn from torchvision import models class EncoderCNN(nn.Module): def __init__(self, d_model=256): super().__init__() resnet = models.resnet18(weights=models.ResNet18_Weights.IMAGENET1K_V1) self.features = nn.Sequential(*list(resnet.children())[:-2]) self.proj = nn.Conv2d(512, d_model, kernel_size=1) def forward(self, x): feat = self.features(x) # (B, 512, 7, 7) feat = self.proj(feat) # (B, d_model, 7, 7) B, C, h, w = feat.shape feat = feat.permute(0, 2, 3, 1).reshape(B, h * w, C) return feat # (B, 49, d_model)

features去掉了 ResNet 最后的全局平均池化和全连接层,保留完整的卷积栈。输入 x 的通道数必须是 3,灰度图在 transform 阶段已经convert('RGB')处理好。如果以后想把特征分辨率提到 14×14,只需要在resnet.children()的截断位置往前挪一个 stage,并把permute后面的 reshape 改成(B, -1, C),后续代码不用动。

解码器用 LSTM 加 Bahdanau 注意力:

class DecoderLSTM(nn.Module): def __init__(self, vocab_size, d_model=256, num_layers=2): super().__init__() self.embedding = nn.Embedding(vocab_size, d_model) self.lstm = nn.LSTM(d_model, d_model, num_layers) self.attn = nn.Linear(d_model * 2, 1) self.out = nn.Linear(d_model * 2, vocab_size) def forward_step(self, prev_token, hidden, enc_feat, mask): emb = self.embedding(prev_token) # (B, d_model) lstm_out, hidden = self.lstm(emb.unsqueeze(0), hidden) lstm_out = lstm_out.squeeze(0) # (B, d_model) N = enc_feat.size(1) q = lstm_out.unsqueeze(1).expand(-1, N, -1) # (B, N, d_model) score = self.attn(torch.cat([q, enc_feat], dim=-1)).squeeze(-1) score = score.masked_fill(mask == 0, float('-inf')) alpha = torch.softmax(score, dim=-1) # (B, N) context = (alpha.unsqueeze(-1) * enc_feat).sum(dim=1) logits = self.out(torch.cat([lstm_out, context], dim=-1)) return logits, hidden, alpha def forward(self, tgt, enc_feat, mask, teacher_forcing=0.5): prev_token = tgt[:, 0] hidden = None logits_list = [] for t in range(1, tgt.size(1)): logits, hidden, _ = self.forward_step(prev_token, hidden, enc_feat, mask) logits_list.append(logits) if torch.rand(1).item() < teacher_forcing: prev_token = tgt[:, t] else: prev_token = logits.argmax(dim=-1) return torch.stack(logits_list, dim=1) # (B, T-1, V)

forward_step是单步计算,输入上一个 token、LSTM 状态、编码特征和掩码。注意力分数先经过masked_fill,把 padding 位置的分数置为负无穷,softmax 后权重自然为 0。forward是训练入口,循环里按 teacher forcing 概率决定下一步输入是标签还是模型自己的预测,这就是计划调度的雏形。hidden 初始为 None,LSTM 内部会默认初始化为零状态,不需要额外处理。

3.4 训练参数表与最小训练循环

有了 tokenizer、数据集和模型,最小训练循环不到 30 行:

import torch import torch.nn as nn from torch.utils.data import DataLoader device = torch.device('cuda' if torch.cuda.is_available() else 'cpu') encoder = EncoderCNN(d_model=256).to(device) decoder = DecoderLSTM(vocab_size=len(tokenizer.char2idx), d_model=256).to(device) params = list(encoder.parameters()) + list(decoder.parameters()) optimizer = torch.optim.Adam(params, lr=1e-3) criterion = nn.CrossEntropyLoss(ignore_index=0) # 0 是 [PAD] 的 id def collate_fn(batch): imgs, seqs = zip(*batch) imgs = torch.stack(imgs) max_len = max(len(s) for s in seqs) padded = torch.zeros(len(seqs), max_len, dtype=torch.long) for i, s in enumerate(seqs): padded[i, :len(s)] = torch.tensor(s) return imgs, padded loader = DataLoader(ds, batch_size=32, shuffle=True, collate_fn=collate_fn) for epoch in range(30): total_loss = 0.0 for imgs, tgt in loader: imgs = imgs.to(device) tgt = tgt.to(device) enc_feat = encoder(imgs) # (B, 49, d_model) B, N, _ = enc_feat.shape mask = torch.ones(B, N, device=device) # 无 padding 的图像序列 logits = decoder(tgt[:, :-1], enc_feat, mask, teacher_forcing=0.5) loss = criterion( logits.reshape(-1, logits.size(-1)), tgt[:, 1:].reshape(-1) ) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(params, 1.0) optimizer.step() total_loss += loss.item() print(f"epoch {epoch}: loss {total_loss / len(loader):.4f}")

代码里几个参数值得单独说。batch_size 32 在 8G 显存附近是安全值;如果只有 6G,降到 16,同时用半精度训练把损失打回去。学习率 1e-3 是 Adam 的常用起点,公式识别任务里 loss 通常第一个 epoch 从 5 左右快速降到 2 附近,如果 5 个 epoch 还没明显下降,优先检查 tokenizer 是否把 LaTeX 符号切散了,其次才是调 lr。clip_grad_norm 取 1.0,LSTM 解码器容易梯度爆炸,不加这个 loss 会时不时跳到 20 以上。下表是这组代码的常用参数组合,可以直接抄进实验记录:

参数推荐值说明
image_size224×224太小丢失上下标,太大增加显存
batch_size328G 显存参考值,OOM 就减半
d_model256编码器和解码器统一宽度
num_layers2LSTM 深度,3 层可尝试但容易过拟合
teacher_forcing0.5→0.0随 epoch 线性衰减,最后让模型自解码
lr1e-3 起步,epoch 20 后降为 1e-4余弦退火更好但初赛阶段不建议
max_len128训练公式 token 上限,超出截断
epochs30是否早停看验证集公式级识别率

4. 公式识别训练与推理的实战参数和排错:稳定复现才是核心

4.1 推理:贪心解码与束搜索的取舍

训练结束后进入推理。贪心解码最简单:每步取概率最高的 token 作为下一步输入,循环到[EOS]。速度快但有两个问题:某个位置选错,后面 token 全被带歪;相邻两步之间概率比较振荡,整条序列的联合分数其实很低。束搜索在每步保留 beam_size 条得分最高的候选序列,最终从候选中选总分最高的那条作为结果。

def beam_search_decode(encoder, decoder, image, beam_size=5, max_len=128): enc_feat = encoder(image.unsqueeze(0)) # (1, N, d_model) B, N, D = enc_feat.shape mask = torch.ones(B, N, device=image.device) beams = [([tokenizer.char2idx['[SOS]']], None, 0.0)] for _ in range(max_len): candidates = [] for seq, hidden, score in beams: if seq[-1] == tokenizer.char2idx['[EOS]']: candidates.append((seq, hidden, score)) continue prev = torch.tensor([seq[-1]], device=image.device) logits, hidden, _ = decoder.forward_step(prev, hidden, enc_feat, mask) log_probs = torch.log_softmax(logits[0], dim=-1) topk = log_probs.topk(beam_size) for idx, val in zip(topk.indices, topk.values): candidates.append((seq + [idx.item()], hidden, score + val.item())) beams = sorted(candidates, key=lambda b: b[2], reverse=True)[:beam_size] if all(b[0][-1] == tokenizer.char2idx['[EOS]'] for b in beams): break best = max(beams, key=lambda b: b[2]) return tokenizer.decode(best[0][1:-1])

参数上 beam_size 取 5 是一个平衡点:beam 1 就是贪心,beam 10 以上推理时间线性上涨但识别率提升很小。代码里每条候选都携带自己的 LSTM hidden,互不干扰;同一轮循环把 top-k 展开到候选列表后统一排序裁剪。如果显存紧张,也可以把 hidden 放 CPU,或者限制最大生成长度。

4.2 评估脚本:编辑距离、公式级识别率与可视化输出

评估时把测试集每条图片用束搜索解码,与真实 LaTeX 比较。编辑距离计算可以直接用动态规划:

def edit_distance(pred, target): m, n = len(pred), len(target) dp = [[0] * (n + 1) for _ in range(m + 1)] for i in range(m + 1): dp[i][0] = i for j in range(n + 1): dp[0][j] = j for i in range(1, m + 1): for j in range(1, n + 1): cost = 0 if pred[i - 1] == target[j - 1] else 1 dp[i][j] = min(dp[i - 1][j] + 1, dp[i][j - 1] + 1, dp[i - 1][j - 1] + cost) return dp[m][n]

评估输出建议保留三列:图片名、预测 LaTeX、真实 LaTeX。把编辑距离按大小排序,排前面的就是最值得人工检查的错误样本。公式级识别率则单独统计完全一致的条数除以总条数,两个指标分开记录。如果一个样本的编辑距离只有 1,但公式级识别率算错,这类结果在文档里要专门讨论,评审喜欢看这种细粒度分析。

4.3 常见坑 1:预训练权重和输入通道不匹配

这是新手踩得最多的一个问题。ImageNet 的预训练卷积接受三通道图,输入要减去 ImageNet 的 mean 除以 std。公式图像如果从灰度图复制成三通道,先检查通道维是不是 (B, 3, H, W),再用transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])归一化。很多模型不收敛,问题不在结构,而是输入张量值域全在 0 到 255 没归一化,梯度在第一个卷积层就异常放大。

4.4 常见坑 2:OOV 与低频符号的欠表达

LaTeX 符号分布极不均衡:x1+出现几千次,\heartsuit可能在全部训练集只出现几十次。字符级 tokenizer 对冷僻符号很敏感,模型大概率把它们编码成[UNK],导致这类公式识别率极低。解决办法:一是统计词频后把低频字符映射到[UNK],并在文档里说明模型无法识别生僻符号;二是扩充对应符号的合成数据,让每个符号至少出现 50 次。实验记录里单独画一条[UNK]占比下降曲线,是答辩时很实在的数据。

4.5 常见坑 3:长度泛化与 max_len 截断

训练时 max_len 设为 128,公式一旦超长直接截断,模型的注意力会丢失后半部分内容。真实场景里公式长度分布和训练集通常不一致,长公式的错误率会从 10% 跳到 50%。缓解办法:训练时按比例混入超长样本,比如 5% 的批次使用 160 的上限;推理时允许 max_len 设为训练上限的 1.5 倍,并在日志里记录被截断的样本,避免静默丢符号。

4.6 常见坑 4:显存不足与 loss 不降

显存不足时除了调小 batch,也可以开启 FP16 混合精度训练,batch size 能翻倍。loss 不降时先排除三个原因:学习率过大、词表把\frac拆散、标签对齐错位。下面这个排查表可以直接贴在实验文档里:

现象优先检查项常见原因
loss 不降归一化、词表切分输入没做 ImageNet 标准化
loss 震荡学习率、梯度裁剪lr 过大或 lstm 梯度爆炸
长公式全错max_len、注意力图训练长度覆盖不足
显存 OOMbatch_size、混合精度特征分辨率太高

5. 把源码和文档说明做成高分毕业设计:目录组织与一页图实验

5.1 源码目录怎么布置

一份交付给评审的源码压缩包,最忌讳的是训练脚本和实验结果全堆在根目录。我一般会这样组织:

formula-ocr/ ├── data/ │ ├── train/ # 渲染好的公式图片 │ ├── train_labels.txt # 文件名<TAB>LaTeX │ ├── test/ │ └── test_labels.txt ├── src/ │ ├── tokenizer.py # 词表构建与编解码 │ ├── dataset.py # Dataset 和 collate_fn │ ├── encoder.py # ResNet 编码器 │ ├── decoder.py # LSTM + 注意力解码器 │ ├── train.py # 训练循环,参数走 argparse │ ├── evaluate.py # 编辑距离/公式级识别率 │ └── inference.py # 单张图束搜索推理 ├── weights/ │ └── best_model.pt # 最优 checkpoint ├── requirements.txt └── README.md # 环境、数据、训练/推理命令

train.py 的训练参数要支持命令行传入,评审如果在自己电脑上复现,改参数比改代码体验好得多。README 写清 Python 3.8 以上、PyTorch 安装方式和两条最小命令:一条训练、一条推理。requirements.txt 不要锁死版本号,只写大版本区间,避免环境冲突。

5.2 文档说明写什么:从模型图到错误分析

文档说明至少要有四块:任务定义加相关工作一段带过,重点写自己的模型结构设计;实验部分交代数据集规模和划分方式;然后是用例分析,放出实际推理结果对比,比如“预测是\frac{1}{2},标准答案是\frac{1}{x}”这类差距;最后是错误分析,给出编辑距离最高的样本截图和人工判定错误类型,比如结构识别错还是符号漏检。这四块写下来,答辩老师能快速看到工作量,评审也不会觉得只是一次普通模型复现。

5.3 一页图技巧:用注意力热力图说明模型“看懂”了公式

最后给一个能直接放进论文的成品技巧:把推理时的注意力权重画成热力图,叠加在原图上,直观展示解码器生成每个 token 时在看图像哪个区域。

import matplotlib.pyplot as plt from torch.nn.functional import interpolate def visualize_attention(image, tokens, attn_weights, feat_h=7, feat_w=7): # image: (3, 224, 224),attn_weights: (T, 49) img = image.permute(1, 2, 0).numpy() fig, axes = plt.subplots(1, len(tokens), figsize=(2 * len(tokens), 2.2)) for ax, token, attn in zip(axes, tokens, attn_weights): ax.imshow(img, cmap='gray') attn_map = attn.reshape(feat_h, feat_w) attn_up = interpolate(attn_map[None, None, :, :].float(), size=(224, 224), mode='bilinear') ax.imshow(attn_up.squeeze().numpy(), cmap='jet', alpha=0.5) ax.set_title(token, fontsize=10) ax.axis('off') plt.tight_layout() plt.savefig('attention_grid.png', dpi=150)

feat_h=7, feat_w=7是 ResNet-18 在 224 输入下的特征网格;如果改动过 stride,这里要同步改。图像保持和输入模型前相同通道顺序。把这张图放到文档的模型分析章节,比任何文字都能说明问题:某个字符的热力图聚焦在图像正确区域,说明对齐机制学到位了;如果注意力散成大片,就直接对应前面提到的长度泛化问题。答辩时评审问“你的注意力真的有用吗”,直接翻这一页图。

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

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

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

立即咨询