☰
基于Transformer的手写文本识别:从数据到部署的工程实践
2026/10/11 19:34:27 网站建设 项目流程

简介:基于Transformer的手写文本识别系统实现与源码解析资源包,面向自然语言处理与计算机视觉交叉方向的研究者和中级开发者,提供一套不依赖字符分割的端到端序列识别方案。方案以编码器-解码器为骨架,通过多头自注意力捕捉字形长程依赖,并融合卷积特征提取、二维相对位置编码与课程学习训练策略;数据预处理涵盖弹性形变增强与笔画归一化,可提升对书写变异的鲁棒性。资源共18个文件,以9个Python源码模块为主,辅以Notebook演示、模型备份与说明文档,压缩包仅132KB;源码模块按数据生成、预处理、模型构建、评估推理划分,结构清晰,便于按需阅读与二次开发。已有133人学习。内容涵盖完整训练流水线、超参数配置、推理部署接口;模型在IAM与CASIA-HWDB上分别达到94.7%和91.2%行级准确率,较LSTM-CTC错误率降低23.6%,在连笔字和倾斜文本上优势明显。适合作为Transformer序列识别项目落地与源码研读的参考。

1. 手写文本识别不是印刷体 OCR 换个大模型就完事:基于 Transformer 的方案到底解决什么问题

把印刷体 OCR 里跑得很好的识别管线原封不动挪到手写文本上,第一轮测试就会翻车——行高不齐、字迹倾斜、连笔、涂抹,这些在清晰印刷体上很少出现的情况,到了手写场景全变成了常态。基于 Transformer 的手写文本识别系统,核心不是把某个视觉模型换成 Transformer 就万事大吉,而是要把「图像特征提取」和「序列依赖建模」用一个统一的框架重新组织起来,让模型既能看清单个字符的形状,又能把握整行甚至整段的书写语境。这套方案适合做票据识别、档案数字化、作业批改这类需要处理自然手写笔迹的工程场景,也适合想读懂 HTR(Handwritten Text Recognition)类开源项目源码、弄清每个模块为什么这么设计的从业者。它的代价同样明显:数据需求、显存占用、训练时间都比传统 CRNN 方案高一截,值不值得上 Transformer,要看你的数据量和场景复杂度。

2. 训练数据怎么准备:数据集选型、文本行切分和标注对齐

2.1 三个公开手写数据集怎么选:先想清楚你要识别谁的字

手写文本识别跟印刷体 OCR 最大的差异在数据分布。印刷体字符形状高度一致,公开数据集随便拉一个就能训出不错的基线;手写则要看字迹风格、语言、版式。常见做法是先确定目标场景,再选数据集,不要一上来就找最大的那个。

  • IAM:英文手写,收录了数百个写手的文本行样本,带单词级和行级标注,是英文 HTR 最常用的基准。做英文票据、信函识别选它最稳。
  • RIMES:法文手写,同样以文本行为单位,包含真实邮件和行政文书,字符分布偏手写体连笔风格,法文场景用。
  • CASIA-HWDB:中文手写,包含离线单字和文本行两套,单字数据集规模大,文本行数据集标注质量较好。做中文档案、表单识别基本绕不开它。

选择时我一般看三个维度:语种是否匹配、标注单位是行还是单字、写手覆盖面是否够广。单字模型容易做,但整行识别时字符切分本身就很难,所以业界普遍直接做「图像到文本序列」的端到端建模,数据集也优先选文本行标注。

提示:这些公开数据集的授权方式和使用限制各不相同,用于商业项目前要逐个确认,不要直接在代码里写死下载地址。

2.2 预处理与文本行对齐:从灰度归一化到固定高度裁剪

拿到文本行图像后,预处理流程看起来简单,实际坑不少。第一步是灰度化和去背景,手写扫描件常有浅色格线、墨渍、纸张纹理,直接用彩色图会增加没必要学的噪声。第二步是最关键的「高度归一化」:把所有图像缩放到同一个高度,典型值是 32 或 48 像素,宽度按比例缩放。保持原始宽高比这一点很重要,如果强行 resize 到固定宽高,字符会被压扁,Transformer 学到的位置编码会变得非常奇怪。

归一化后还要处理一个常见问题:同一批数据里文本行长度差别很大。有的行 30 个字符,有的行 120 个字符,直接组成 batch 会让短的被 pad 到很长,Transformer 的 attention 计算量随序列长度平方增长,大部分算力都浪费在 pad 上。常见做法是训练时按长度分桶(bucket),推理时动态调整最大宽度。

def preprocess_line(img: np.ndarray, target_h: int = 48) -> tuple[np.ndarray, float]: """ 文本行预处理:灰度化 -> 去背景 -> 等比缩放到固定高度 返回处理后的图像和缩放比例,缩放比例用于后续把预测框映射回原图。 """ if img.ndim == 3: gray = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY) else: gray = img # 去背景:大卷积核的形态学操作去掉浅色纹理(格线、纸纹) blur = cv2.GaussianBlur(gray, (5, 5), 0) thresh = cv2.adaptiveThreshold( blur, 255, cv2.ADAPTIVE_THRESH_GAUSSIAN_C, cv2.THRESH_BINARY, 31, 15 ) # adaptiveThreshold 的输出是黑白图,但训练时不要做二值化, # 保留灰度梯度能让模型对扫描亮度的变化更鲁棒。 gray = cv2.bitwise_and(gray, gray, mask=255 - thresh) h, w = gray.shape scale = target_h / h new_w = int(round(w * scale)) # 宽度限制为 8 的倍数,方便后续下采样和 attention padding new_w = max(8, int(new_w / 8) * 8) resized = cv2.resize(gray, (new_w, target_h), interpolation=cv2.INTER_AREA) # 转为 0~1 的 float32,并做均值方差归一化 resized = resized.astype(np.float32) / 255.0 return resized, scale

这段代码里值得注意的点是adaptiveThreshold的部分:它的输出并不是直接作为训练图像,而是用来生成 mask 过滤掉浅色背景。如果直接把二值化后的图送进模型,字迹的墨色深浅信息会丢失,模型对铅笔、颜色较浅的圆珠笔笔迹会非常脆弱。缩放到固定高度而不是固定尺寸,是 HTR 任务和普通 OCR 的明显区别,Transformer encoder 端的序列长度等于图像宽度下采样后的 token 数,保持宽高比才能让位置编码学到「字符在宽度方向上的先后关系」。

2.3 DataLoader 实现:标签编码、padding mask 和弹性数据增强

预处理之后就是 DataLoader 的活。手写识别的标签是字符串,比如"Dear Sir, I am writing to...",需要转成字符索引序列。这里有个容易搞错的地方:CTC 和 Transformer decoder 的标签编码方式不同。CTC 需要额外预留blank索引(通常放最后),而自回归 decoder 需要<sos>和<eos>两个特殊 token。如果代码里把这两套特殊 token 混在一起,训练时 loss 会诡异跳变。

数据增强方面,手写文本适合做轻度弹性变形(elastic distortion),这是模拟真实书写时纸张起伏和笔画抖动的常用手段,比旋转裁剪更有效。随机亮度抖动和局部擦除(random erasing)也很有用,能模拟墨迹断线和涂抹遮挡。

class HTRDataset(torch.utils.data.Dataset): def __init__(self, samples, char_to_idx, target_h=48, augment=True): """ samples: list of (image_path, text) char_to_idx: 字符到索引的映射,需包含 blank/ sos / eos 三个特殊token """ self.samples = samples self.char_to_idx = char_to_idx self.target_h = target_h self.augment = augment def __getitem__(self, idx): img_path, text = self.samples[idx] img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) img, _ = preprocess_line(img, self.target_h) if self.augment: # 轻度弹性变形,网格点扰动范围控制在 1~2 像素 img = elastic_distort(img, alpha=3, sigma=0.08) # 随机水平方向轻微斜切,模拟书写倾斜 img = random_shear(img, max_angle=0.15) # 标签编码:注意这里不包含空格填充,pad 在 collate_fn 里统一做 chars = list(text) label = [self.char_to_idx[c] for c in chars if c in self.char_to_idx] label_tensor = torch.tensor(label, dtype=torch.long) return torch.from_numpy(img).unsqueeze(0).float(), label_tensor

elastic_distort和random_shear的具体实现可以用 OpenCV 的remap做网格插值,参数不要给太大。alpha 超过 5 会把字符形状扭曲到难以辨认,反而拉低训练效果。这套增强对 Transformer 类模型尤其重要——它比 LSTM 更依赖数据多样性来学到稳定的字符形状表征,小数据集上不增强,attention 很容易只记住几个高频字形。

3. Transformer 模型架构怎么搭:特征提取、位置编码和解码头选型

3.1 为什么手写场景值得上 Transformer Encoder:从 CRNN 到视觉 Transformer 的演进逻辑

早期手写识别的主流结构是 CNN + BiLSTM + CTC:CNN 下采样出特征图,BiLSTM 按时间步建模左右上下文,CTC 做序列对齐。这套结构在规整手写体上效果不错,但有两个硬伤:一是 BiLSTM 是串行的,训练慢,长文本行容易遗忘早先的字符信息;二是在二维结构上它只能沿水平方向传递信息,遇到抬头、换行、字符间距异常时就只能靠 CNN 的局部感受野硬撑。

基于 Transformer 的手写文本识别系统,最常见做法是把 BiLSTM 替换为 Transformer Encoder,CNN 部分可以保留也可以整个换成 ViT 风格的分块嵌入。Transformer Encoder 的 self-attention 能直接建模「当前字符与整行其他字符」的关系,不会因为序列太长而丢失远处信息。更关键的是,它可以并行训练,GPU 利用率比 BiLSTM 高得多。代价是 attention 的显存占用随序列长度平方增长,而且需要更多的数据来拟合,这个 trade-off 在选型时就要想清楚。

这里要区分一个概念:视觉 Transformer(ViT)直接对图像分块,而 HTR 场景更多采用「CNN 下采样 + Transformer 编码器」的混合结构。原因很实际——手写字符是细长形状,ViT 的固定分块粒度不容易同时兼顾小写字母的细节和长单词的整体结构,CNN 先把手写文本行缩成兼顾高度和宽度信息的特征序列,Transformer 再在这个序列上建模,效果通常更稳。

3.2 位置编码是手写识别里最容易忽视的配置项

Transformer 本身没有顺序概念,位置编码就是给每个 token 注入「它是第几个字符」的线索。印刷体 OCR 常用标准正弦位置编码,因为字符间距均匀;手写文本的字符宽度差异大,同一行里有窄的i、l和宽的m、W,绝对位置并不能准确反映字符边界。

我的做法是:CNN 特征图的高度维度已经包含了字符的纵向结构信息,所以主要对宽度方向计算位置编码;但二维可学习位置编码在复杂版式中更稳,它会同时把「第几行、第几列」的信息编码进去。实现上,把 CNN 输出的特征图B, C, H, W压缩成B, L, C(L = H * W),然后加上一个可学习的绝对位置嵌入。这么做比正弦位置编码灵活,模型能自己学会不同写手的书写偏移。

class PositionalEncoding(nn.Module): def __init__(self, d_model: int, max_len: int = 512): super().__init__() # 可学习位置编码,而非固定正弦编码 self.pos_embed = nn.Parameter(torch.zeros(1, max_len, d_model)) nn.init.trunc_normal_(self.pos_embed, std=0.02) def forward(self, x: torch.Tensor) -> torch.Tensor: """ x: (B, L, d_model) 只取前 x.size(1) 个位置向量,超出 max_len 部分会被截断。 """ return x + self.pos_embed[:, : x.size(1), :]

注意nn.init.trunc_normal_初始化位置编码很关键,如果用零初始化会让模型前期训练非常慢。我在几个项目里对比过,可学习位置编码在同尺寸模型上比正弦编码带来 3%~5% 的 CER 下降,代价是训练时要保证文本行宽度不超出max_len,否则超出部分没有位置向量可用。所以 DataLoader 里的bucket策略作用不只是省显存,还直接决定了位置编码的有效覆盖范围。

3.3 Encoder 与解码头选型:CTC 还是自回归 Attention

模型骨架确定后,解码头是第二个决策点。目前主流有两种路线:一种是在 Transformer Encoder 末端接线性层 + CTC Loss;另一种是接一个轻量的自回归 Transformer Decoder,逐个字符预测。CTC 方案实现简单、推理速度快,训练时能容忍标签与特征序列长度不对齐;自回归方案理论上能建模字符间的上下文依赖(比如 "qu" 后面几乎一定是元音),但训练和推理都更重,容易在小数据集上过拟合。

实际项目中我做这样的取舍:数据集小于 5 万行文本时,优先选 CTC 头;数据量充足且目标场景有强语言规律(如医疗处方、法务文书),自回归 Decoder 带来的收益才值得额外付出的训练成本。CTC 头还有一个优势是它天然与 CNN 特征图的序列长度兼容,不需要额外设计start/endtoken 的对齐逻辑。

class HTREncoder(nn.Module): def __init__(self, d_model=256, nhead=8, num_layers=6, num_chars=128): super().__init__() # CNN backbone:三层下采样,把高度 48 缩到 12 self.cnn = nn.Sequential( nn.Conv2d(1, 64, kernel_size=3, padding=1), nn.ReLU(), nn.MaxPool2d((2, 2)), nn.Conv2d(64, 128, kernel_size=3, padding=1), nn.ReLU(), nn.MaxPool2d((2, 2)), nn.Conv2d(128, 256, kernel_size=3, padding=1), nn.ReLU(), nn.MaxPool2d((2, 1)), # 高度向下采样,宽度尽量保留 ) self.proj = nn.Conv2d(256, d_model, kernel_size=1) self.pos_embed = PositionalEncoding(d_model, max_len=256) encoder_layer = nn.TransformerEncoderLayer( d_model=d_model, nhead=nhead, dim_feedforward=1024, dropout=0.1, activation="gelu", batch_first=True, ) self.encoder = nn.TransformerEncoder(encoder_layer, num_layers=num_layers) # 映射到字符表大小,CTC 训练时用 self.fc = nn.Linear(d_model, num_chars) def forward(self, x): # x: (B, 1, H, W) feat = self.cnn(x) # (B, 256, H', W') b, c, h, w = feat.shape feat = self.proj(feat) # 压缩通道到 d_model feat = feat.permute(0, 3, 2, 1).reshape(b, h * w, c) # (B, L, d_model) feat = self.pos_embed(feat) feat = self.encoder(feat) # (B, L, d_model) logits = self.fc(feat) # (B, L, num_chars) return logits.permute(1, 0, 2) # CTC 需要 (L, B, num_chars)

这段代码里的MaxPool2d((2, 1))是刻意为之:高度缩半,宽度不缩,目的是让 CNN 输出的 token 序列仍然保持足够高的水平分辨率。如果宽度方向也做 2 倍下采样,输入宽度只有 32 像素的短文本行会只剩 16 个 token,模型几乎无法完成对 20 多个字符的序列建模。这里还有一个细节——CTC 头的输出长度 L 必须大于等于标签长度,所以 CNN 阶段不能过度压缩宽度。输入图像高度 48、宽度 160 时,三层卷积后输出长度大约 640 左右(高度下采样 8 倍,宽度下采样 2 倍),这个比例对大多数文本行都够用。

4. 训练与源码结构解析:损失函数、学习率策略和工程代码组织

4.1 CTC Loss 和 CER 的对应关系:为什么 loss 降了指标没动

训练阶段的核心指标是字符错误率(CER),但 loss 用的是 CTC Loss,这两个指标之间不是严格单调关系。CTC Loss 输出的是整条序列所有可能对齐路径的负对数似然总和,它关注的是「模型整体概率是否提升」;CER 则是在 greedy 解码后逐字比较删除、插入、替换错误。经常出现 loss 稳步下降、CER 却卡住不动的情况,这多半是解码策略的问题而不是模型没收敛。

解码时最容易出错的地方是重复字符处理。CTC 的 blank 机制会把"book"预测成"bok"或"b ook",因为相同字符连续出现时,相邻帧的重复输出会被 blank 隔开后折叠。如果你的标签里大量出现双写字母、叠字,一定要在解码后做一个简单的后处理:把连续相同且中间无 blank 的字符合并。

def greedy_decode(logits, char_to_idx, blank_idx): """ logits: (L, B, num_chars) 模型输出 返回: list[str] 每个样本的解码文本 """ preds = logits.argmax(dim=-1) # 每帧取概率最大的字符索引 id2char = {v: k for k, v in char_to_idx.items()} results = [] for b in range(preds.shape[1]): chars = [] prev = None for t in range(preds.shape[0]): idx = preds[t, b].item() if idx != blank_idx and idx != prev: chars.append(id2char[idx]) prev = idx results.append("".join(chars)) return results

这套贪心解码是基线,项目上线前建议换 beam search 或带语言模型重打分。注意prev变量的作用:它记住上一帧的字符,只有当当前帧不是 blank 且和上一帧不同时才输出,这就是 CTC 的「折叠重复」逻辑。如果用 transformers 库自带的generate方法走的是自回归解码,完全另一条逻辑,不要混用。

CER 计算我建议用jiwer库的cer函数,它内部实现了标准编辑距离。需要特别向项目组说明的是:CER 不区分大小写时,需要在计算指标前把标签和预测都做一个大小写归一化,否则英文文本行会因为首字母大写产生高错误率假象。

4.2 训练超参数表:这些参数是手写识别任务的经验起点

手写识别模型的超参数敏感度很高,我把它按「必须先确定的和翻车后再调的」分成两类。下面的参数配置基于「高度归一化到 48、CNN + Transformer Encoder + CTC」的常见路线,数据量在 10 万~50 万行之间。

参数常见取值说明
d_model256(小) / 384(中)小于 128 时 attention 表达能力不足,大于 512 时小数据过拟合明显
num_layers4 ~ 86 层是一个性价比拐点,再加深对 CER 的改善通常在 0.5% 以内
nhead8必须能被d_model整除
warmup_ratio0.1前 10% 的 step 线性从 0 升到峰值
peak_lr5e-4(AdamW)配合 batch_size 32;batch 翻倍则 lr 相应开根号上调
dropout0.1 ~ 0.2数据集小于 5 万行建议 0.2
label_smoothing0.0(CTC)CTC 头一般不用标签平滑,自回归头可用 0.1
max_length256超出此宽度直接丢弃或在 Dataloader 里过滤
grad_clip1.0手写任务梯度爆炸常见,clip 必须开
batch_size16 ~ 64以显存不爆、显存利用率 > 50% 为准

这些参数不是从论文里抄的,而是几个项目里跑出来的「起点值」。真实调参时,我先把batch_size顶到显存上限的 80%,然后按 batch 大小反推peak_lr,再固定warmup_ratio跑 10 个 epoch 看 loss 曲线形态。如果前期 loss 下降很猛后期震荡大,说明 lr 偏高;如果 loss 一直在高位平着走,先查标签编码再做数据可视化,不要急着调参。

4.3 源码解析:从 config 到 inference 的目录结构与核心链路

既然标题里带「源码解析」,这里用一个常见的 HTR 项目结构来说明各文件职责,这样你拿到任何开源项目都能快速定位关键逻辑。

htr_project/ ├── config.py # 所有超参数和路径配置,用 dataclass 集中管理 ├── dataset.py # 数据集类、数据增强、标签编码 ├── model.py # CNN + Transformer Encoder + 解码头 ├── trainer.py # 训练循环、验证逻辑、checkpoint 保存 ├── decode.py # greedy / beam search 解码 ├── evaluate.py # CER、可视化工具 └── inference.py # 加载模型做单张图片推理

最值得细读的是trainer.py。它的核心链路通常是:从DataLoader取 batch → 前向得到 logits → 计算 CTC Loss → 反向传播 → 梯度裁剪 → 优化器 step → warmup 调度器 step。我建议重点关注三个细节:一是loss是否对log_probs做了log_softmax,CTC Loss 的输入要求是对数概率,很多人直接传 softmax 后的结果导致 loss 为负或者 nan;二是验证时是否在torch.no_grad()和model.eval()模式下关闭了 dropout——Transformer 的 dropout 在推理时一旦忘了关,CER 会显著变差;三是 checkpoint 存储时除了 model state dict,还要保存char_to_idx的副本,否则换机器推理时字符表不匹配会产出完全乱码的文本。

def train_one_epoch(model, dataloader, optimizer, scheduler, criterion, device): model.train() total_loss = 0 for images, labels in dataloader: images = images.to(device) labels = labels.to(device) # 注意:labels 直接给 CTC,不需要 decoder 的 shift 逻辑 logits = model(images) # (L, B, num_chars) log_probs = nn.functional.log_softmax(logits, dim=-1) input_lengths = torch.full( (images.size(0),), logits.size(0), dtype=torch.long, device=device ) target_lengths = (labels != pad_idx).sum(dim=1) loss = criterion(log_probs, labels, input_lengths, target_lengths) optimizer.zero_grad() loss.backward() # 梯度裁剪对 transformer 是必选项 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() scheduler.step() total_loss += loss.item() return total_loss / len(dataloader)

CTC Loss 的input_lengths容易写错:它等于模型输出的时间步长度,也就是logits.size(0),不是图像的原始宽度,也不是标签长度。如果你错误地把input_lengths设成了 batch 里每条样本的真实长度,损失函数会下标越界或报AssertionError,这是 Transformer + CTC 方案最经典的翻车点之一。另外这个代码里的labels是 padding 过的 2D 张量,target_lengths需要去掉 padding 符号再统计,这两处只在处理变长 batch 时才需要关注,固定长度 batch 可以简化。

5. 容易翻车的 5 个地方:从数据标注到解码,逐条排查

5.1 训练 loss 不降反升,第一嫌疑人永远是标签对齐错误

现象:loss 曲线在十几个 epoch 里没有任何下降趋势,甚至偶发跳到数倍正常值(比如从 80 跳到 200),CER 一直在 95% 以上,模型输出在字符表里完全随机。原因多半不是模型,而是 DataLoader 返回的(图像, 标签)配对是乱的。手写数据集打包时经常有「每行图像一个目录,标签在某个 CSV 里按行号索引」的情况,代码里一旦用了os.listdir自带的无序排序,图像和文本就错位了。

解决:在 dataset 初始化时打印前 5 对样本,人工核对图像的字符数和标签的字符数是否大致对应。再进一步,写一个快速自检:把每张图像缩到 32x8 的极低分辨率,和它的标签一起输出到一个 HTML 页面里,肉眼扫一眼。这招看起来笨,但能救回半天排查时间。

5.2 Transformer 在小数据集上过拟合,dropout 和 warmup 怎么配合

现象:loss 在训练集上降到很低,CER 甚至逼近 0,但验证集 CER 高于训练集 20 个百分点以上。这是 Transformer 类模型的典型症状——参数量大、注意力机制灵活,几万行数据很容易被模型背下来。原因不只是「数据少」,还包括位置编码学到的绝对位置信息在训练集上被过度利用,换一个书写位置就失效。

解决:先动两个配置,不要一上来就换模型结构。第一,把dropout从 0.1 提到 0.2,同时确认TransformerEncoderLayer里每个子层都配了 dropout;第二,warmup_ratio不要低于 0.1,让模型在前期少接触高学习率下的极端梯度。如果还过拟合,再加一条数据增强:随机裁剪掉文本行左右 5%~10% 的宽度,逼模型不要依赖「字符在第几个绝对位置」来做预测,而要学会看周围的字符关系。

注意:这两招都试完仍过拟合,就应该考虑减小num_layers或d_model而不是无限加数据增强。增强过度会让字迹变形失真,反而破坏语义。

5.3 解码结果总是合并连续重复字符:不是模型的问题,是 CTC 折叠逻辑踩坑

现象:标签是"effort",预测结果变成"efort";标签是"better",预测出"beter"。这类错误在 CER 里占比极高,但不代表模型没学会字符f或t怎么写。

原因:前面 4.1 节代码里已经解释了,CTC 将「同一字符的连续重复输出」折叠为一个字符,只有当两个相同字符间存在 blank 帧时才能区分"ff"和"f"。手写连笔常常让相邻相同字符的笔迹连在一起,模型很难在两个f之间产生一个清晰稳定的 blank 帧,于是折叠时就丢了一个字符。

解决:最有效的不是改解码,而是在训练阶段把"ff"这类标签里的重复字符用一个额外 token 替换(比如"effort" -> "effort"中间插入/),然后在后处理里把这个 token 替换成重复字符。这个技巧在英文手写数据集上能稳定降低 1~2 个百分点的 CER,代价是字符表里多了一个 token。如果你的场景是中文,叠字结构不同,这个办法收益有限,主要靠 beam search 加语言模型约束。

5.4 显存 OOM 和 batch 丢失:动态 padding 和梯度累积怎么配

现象:训练刚开始一切顺利,跑了几个 epoch 后开始报CUDA out of memory;或者 batch 大小设成 32 时显存刚够,换了一批更长的文本行就爆掉。原因很直接:文本行图像宽度变化大,长样本的 token 序列长度可能是短样本的 3 倍,attention 的显存占用随长度平方增长,均匀分布的长样本会把显存峰值推到远超平均值的位置。

解决:第一步,训练时把样本按宽度分桶,短的先训练,长的后训练,避免一个 batch 里出现两个极长样本;第二步,batch_size不要一次给满,先用 16 跑几个 epoch 观察显存峰值,再用gradient_accumulation_steps=2凑回等效 batch 32 的效果。这两步一起做,OOM 基本能消除。还要检查pin_memory=True是否开启——它能让 CPU 数据加载和 GPU 计算重叠,减少显存上的峰值压力。

5.5 attention 可视化看起来很漂亮,但模型实际效果很差的诡异情况

现象:把 attention 权重可视化后,热力图显示模型确实在看字符周围的上下文,但 CER 始终在 30% 以上,每个词都错一点点。

原因:这一步往往不是模型的问题,而是损失函数和评估指标不匹配。CTC Loss 优化的是整条手写文本行的概率,它不关心每个字符的边界是否精确;而 CER 的计算是字符级的精确匹配。模型学会了「大概的字符序列」,但每个字符的位置编码偏差了半格,解码出来就错位了。可视化里看到 attention 有规律并不代表它学到了正确的字符边界。

解决:把手写文本行的识别结果和标签逐字符对齐后输出到日志里,看是「插入错误为主」还是「替换错误为主」。插入错误多,说明解码后处理该加重叠合并;替换错误多,说明字符表映射或者图像预处理有问题。我通常会在验证集上固定输出 20 条样本的(image, label, pred)三元组到 HTML 页面,每行配一张缩略图。解决 after fifty rounds:你一百次盯着 loss 曲线,不如看二十张图来得直观。

6. 验证系统是否真的可用:CER 曲线分析、beam search 重打分和模型导出实践

模型在验证集上基本收敛后,下一步不是直接部署,而是做一次系统级的验证和优化循环。我会先固定一个测试集(和训练集来自不同写手或不同批次的数据),跑出三个数值:greedy CER、beam search CER、带语言模型重打分的 CER。如果后两者比前者有明显改善,说明模型对字符形状的识别能力已经足够,瓶颈在语言上下文;如果三者差距很小,说明模型本身还没吃透字迹特征,应该回头补数据或调架构。

一个可靠的验证流程是:先每 5 个 epoch 存一次 checkpoint,记录对应的 CER,画出 CER 随训练步数的曲线。正常情况应该是先快速下降、再平缓波动;如果曲线在下降一段后又反弹上去,说明开始过拟合,取最低点的 checkpoint 而不是最后一个 checkpoint 部署,这个习惯能救回不少效果。我一般会在训练脚本里默认开启「best model 保存」,并保留完整的 CER 历史——手写文本识别的模型迭代是很吃经验的,保留每一版的结果能让你准确判断改动是否正向,而不是靠「感觉上一版更好」。

部署阶段的技巧是模型导出时不要直接用 PyTorch 的torch.save存整个模型对象,而是导出为 TorchScript 或 ONNX,同时把推理时的预处理函数也一起固化进去。手写文本识别系统最容易在「训练环境能跑、生产环境报错」的边界上出问题——训练脚本里图像缩放用的cv2.resize和生产环境用的参数不一致,或者字符表顺序在导出时被重新洗牌,这些低级错误会导致模型上线后输出一堆乱码。导出的模型里应该包含一个「带预处理管道」的完整前向函数,输入是一张原始图像,输出是字符串,这样从环境隔离的角度看,模型的行为是可复现的。

最后我想说,手写文本识别里「看起来对了」和「真的对了」之间,隔着一整套验证和理解。先把解码头在干净验证集上跑到的 CER 打到 10% 以内,再去谈调参。(当然,没有配 dropot 和 warmup 就训 Transformer 的话,连 30% 都难。)与其纠结是换大模型还是用更深的 encoder,不如先把这套「数据—模型—解码—验证」的闭环里每一个环节都做成可记录、可对比的版本。每调一版,我都会把训练参数、数据集版本、CER 记录在一个简短的日志里——这个习惯比任何单次调参技巧都值钱,希望帮到你。

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

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

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

立即咨询