☰
Transformer重做手写OCR:端到端识别架构与PyTorch实战
2026/10/11 22:54:48 网站建设 项目流程

简介:面向手写文本识别与Transformer序列建模开发者/学生的系统实现与源码解析资源包,覆盖编码器-解码器结构的端到端笔迹识别流程,无需字符分割预处理,尤其适配连笔字、倾斜文本等复杂场景。压缩包共18个文件,以9个Python源码文件为主,配套示例笔记、说明文档、依赖清单与备份文件,整体约132KB,结构紧凑。资源目前已有133人学习。内容覆盖数据弹性形变增强、笔画归一化、卷积特征提取、二维相对位置编码与课程学习训练策略等关键环节,并提供完整训练流水线、超参数配置、推理接口、评估指标体系和可视化分析组件。IAM英文手写数据库上94.7%、CASIA-HWDB中文数据集上91.2%的行级识别准确率可供参考,较传统LSTM-CTC错误率降低23.6%,适合需要理解Transformer在OCR领域落地方法、快速复现实验并扩展到自有手写数据集的读者。

1. 手写文本识别:为什么 Transformer 值得你重做一遍 OCR

你手里有一批历史档案的扫描件,行文是手写体,连笔、歪斜、字间没有固定间隔,丢给传统 OCR 引擎返回的是大片乱码。基于 Transformer 的手写文本识别系统解决的正是这类问题:把一整行手写图像直接映射成文本序列,不需要逐字分割,也不需要人工标注每个字的位置,训练数据只需要“图像 + 整行文本”。这个方向不只在论文里成立,档案数字化、表单审批、医疗记录转写等场景已经跑出了实际价值。适合正在优化 OCR 准确率、或者想把笔迹识别落地成内部服务的工程师阅读。接下来我会从模型选型讲到数据切分,再拆解一套可运行的 PyTorch 实现,最后把训练和部署里最容易翻车的几个坑一次说清。

2. 为什么是 Transformer:先把手写识别架构的账算清楚

手写文本识别和印刷体 OCR 最本质的区别在于“字符边界不可靠”。印刷体可以先检测字符框再分类,手写行里两个字连笔的情况下根本没有干净边界,所以现代方案普遍走“整行识别”路线,也就是输入一整行图像,输出一段字符序列。这样问题就被抽象成:给定 2D 图像特征,生成 1D 字符序列。

早年主流的深度学习方案是 CNN + BLSTM + CTC。BLSTM 沿水平方向记忆上下文,CTC 负责把帧级输出对齐成字符序列。这套组合在公开英文手写集和开源中文手写语料上都拿过不错的分数,但它有两个硬伤:BLSTM 的时间步必须顺序执行,训练和推理都比 CNN 慢一个量级;长文本行上梯度沿时间方向衰减,序列到后半段时历史信息会明显弱化。Transformer 进来以后,自注意力让任意两个位置直接相连,一层的感受野就是整行;训练也能并行,GPU 面对的不再是时间步串行计算。代价是自注意力不天然带顺序性,位置编码必须设计好,否则模型分不清“很多”和“多很”到底哪个是对的。

2.1 RNN 的顺序建模瓶颈与 Transformer 的并行收益

手写行图像的强顺序性体现在多个层面:相邻字符之间笔画有承接,书写风格在整行内保持一致,倾斜角度从左到右缓慢变化。BLSTM 确实能捕捉这些信息,但它的双向结构在工程实现上要跑两遍时间步,训练吞吐量上不去;梯度经过几十个时间步之后,哪怕有 LSTM 的门控,也会出现长距离信息稀释。Transformer 每个位置上都能直接看到全序列,建模长距离笔画依赖的能力来得更直接。

但 Transformer 不是拿来就能用。手写行在视觉上是高瘦图,宽度远大于高度,如果用 ViT 那种方形 patch 切法,一个 patch 可能横跨好几个字,语义被切碎。常见做法是先用 CNN 把图像压成“一列一个特征向量”的序列,再交给 Transformer 编码器。这样每个 token 对应原图里约 8 像素宽的一条竖带,token 之间天然按水平顺序排列,Transformer 只需要在这个顺序上补一套位置编码。工程上这个方案的好处是:CNN 负责视觉细节,Transformer 负责序列上下文,分工明确,调参时哪块出问题直接定位到对应模块。

2.2 解码头选型:Encoder-Decoder 与 Encoder + CTC 的取舍

编码器输出的是一串特征,怎么把它变成文本序列,有两条主流路线。第一条是 Encoder-Decoder,Decoder 端每个时间步生成一个字符,训练时用 teacher forcing。它的优点是建模了字符之间的条件依赖,比如“么”后面更可能是“样”而不是“地”;缺点是训练需要每一步的对齐或者至少强语言信号,数据量不足时容易学出重复生成和漏字的毛病,并且逐字符生成在推理时无法并行。

第二条是 Encoder + CTC。CTC 只要求整行文本标签,路径聚合自动处理“哪一帧对应哪个字符”的分配问题。手写识别场景里,字写歪、连笔、涂改都是常态,CTC 对这类不对齐情况非常宽容。工程上我更倾向 CTC 头,理由很现实:生产环境里最难的不是模型结构,而是标注数据。CTC 只需要“图像 + 整行文本”的现成标注流程,不需要维护字符框,不用逐字对齐。Encoder-Decoder 在英文场景里如果标注质量高、语料充足,表现可能略好,但把它当 baseline 去做业务,CTC 永远是稳的那一个。

2.3 推荐基线结构:CNN Backbone + Transformer Encoder + CTC 头

下面是我在模拟项目X里反复用的一套基线结构,所有参数都可以当作起点,不用一上来就堆大模型。

模块选型理由
CNN 主干ResNet-18/25 去掉最后一个 stage,输出 stride 8保留足够宽度的列特征,连续笔画不糊
序列化自适应池化把高度压成 1,得到 (B, W', C)每个 token 对应原图一条竖带
位置编码可学习绝对位置编码,max_len 512简单可靠,中文行宽 512px 以内够用
Transformer4 层,d_model 256,nhead 8,FFN 1024,Pre-Norm浅层编码器对数据量要求更低
解码头Linear(d_model, num_classes) + log_softmax配合 CTCLoss,不需要额外语言模型
优化器AdamW,峰值学习率 1e-4 到 3e-4,warmup 500 步Transformer 的收敛对学习率十分敏感

在写完整代码之前,先用一段形状验证脚本把这条路走通,避免后面模型搭了半天,输入输出维度对不上。以下代码假定输入图像已经预处理成高度 64、宽度 512、单通道。

import torch # 输入形状模拟:batch=2, 通道=1, 高=64, 宽=512 x = torch.randn(2, 1, 64, 512) # 假设已经构建好三个核心模块 feat = backbone(x) # 输出 (2, 128, 1, 64) feat = feat.squeeze(2) # 去掉高度维,变成 (2, 128, 64) feat = feat.permute(0, 2, 1) # 转成序列模式 (2, 64, 128) feat = pos_enc(proj(feat)) # 线性映射 + 位置编码 (2, 64, 256) out = encoder(feat) # Transformer 编码 (2, 64, 256) logits = ctc_head(out) # 分类头 (2, 64, num_classes)

逻辑说明:宽度 512 经过 8 倍下采样后得到 64 个 token,每个 token 在原始图上覆盖 8 个像素宽的竖条。squeeze(2)丢掉高度维的前提是 CNN 已经把高度压成 1,这一步必须在池化层完成,不能省。permute把通道维挪到序列维后面,因为 Transformer 编码器的batch_first=True约定输入是 (batch, seq_len, d_model)。这段代码能跑通,说明整体维度链路没问题,接下来怎么切数据、怎么喂标签,心里就有底了。

3. 数据管线与预处理:从整页扫描件到干净的文本行样本

模型设计得再好,数据切分和预处理跟不上也是白搭。手写识别场景里,原始数据经常是整页扫描件,而不是干净的文本行图片。第一步要做的就是把页面切成行,这一步做得越干净,后面训练省的心越多。常见的做法是水平投影法:二值化之后沿垂直方向统计每行的暗像素数量,行与行之间有空白间隙,投影值就会掉到接近零,据此切出每行的上下边界。

3.1 页面分割与文本行提取:水平投影法的最小实现

下面这段代码可以处理绝大多数规整的扫描件。它假设文档方向基本水平,如果整页旋转超过两三度,要先做倾斜校正,不然投影切出来的行边界是歪的,字符会被拦腰切断。

import cv2 import numpy as np def extract_lines(image_path, min_h=12, gap=3): img = cv2.imread(image_path, cv2.IMREAD_GRAYSCALE) # 反二值化,让笔画变成白色、背景变成黑色,便于投影统计 _, bin_img = cv2.threshold(img, 0, 255, cv2.THRESH_BINARY_INV + cv2.THRESH_OTSU) # 按行求和:同一行里的笔画像素越多,该行的投影值越大 proj = bin_img.sum(axis=1) / 255.0 # 投影值大于 1.0 的位置视为有笔画的行 idx = np.where(proj > 1.0)[0] if len(idx) == 0: return [] lines = [] start, prev = idx[0], idx[0] for i in idx[1:]: if i - prev > gap: # 行间间隙超过 gap 像素,判定为换行 if prev - start + 1 >= min_h: lines.append((start, prev + 1)) start = i prev = i if prev - start + 1 >= min_h: lines.append((start, prev + 1)) return [img[s:e, :] for s, e in lines]

逻辑说明:THRESH_BINARY_INV + THRESH_OTSU把灰度图变成反色二值图,笔画是白色,统计时每行白色像素越多,投影值越大。gap=3表示两个行块之间如果只有不到 3 个像素的间隔,就视为同一行里的笔画像素抖动;min_h=12过滤掉高度过小的噪声区域,比如下划线、纸张污渍。切出来后,我一般会在每个行块上下各扩 4 个像素再裁图,避免手写字母的最高点和最低点被切掉。英文手写行更扁,min_h可以降到 8;中文手写行高一些,12 到 20 都常见。

3.2 行图归一化与数据增广:高度固定、宽高比保持、扰动克制

拿到文本行图像后,不能直接塞进模型。手写行长短不一,宽高比差距很大,必须先把高度统一。通常的做法是固定高度为 64 或 128,宽度按比例缩放,超出上限的截断,不足的右侧补白。高度选 64 的优点是显存开销小,适合快速验证;生产模型可以上 128,细笔画保留得更完整。

def preprocess_line(img, target_h=64, max_w=512, pad_val=255): # 先扩边,避免紧贴边缘的笔画在缩放时被削掉 img = cv2.copyMakeBorder(img, 4, 4, 4, 4, cv2.BORDER_CONSTANT, value=[pad_val] * 3) h, w = img.shape[:2] scale = target_h / h new_w = int(w * scale) img = cv2.resize(img, (new_w, target_h), interpolation=cv2.INTER_AREA) # 超宽样本直接截断,防止一张图拖垮整个 batch if new_w > max_w: img = img[:, :max_w] new_w = max_w # 右侧补白到 8 的倍数,方便 CNN 阶段的下采样 pad_w = (8 - new_w % 8) % 8 if pad_w > 0: img = cv2.copyMakeBorder(img, 0, 0, 0, pad_w, cv2.BORDER_CONSTANT, value=[pad_val] * 3) # 归一化到 [-1, 1] img = (img.astype(np.float32) / 255.0 - 0.5) / 0.5 return img

逻辑说明:INTER_AREA在缩小图像时能保留笔画的整体结构,不容易出现锯齿伪影。max_w=512是经验值,超过这个宽度要么是扫描分辨率太高,要么是文本行长度超过 30 个汉字,截断后丢掉的信息在数据切分阶段就应该被处理掉,而不是让模型硬学。补白到 8 的倍数是为了配合 CNN 的两次 stride 2 下采样,避免在池化时出现非整除的边界问题。手写输入归一化到 [-1, 1] 而不是 ImageNet 的均值方差,原因是背景白色、笔画黑色,分布比自然图像简单,用全局归一化让模型更快收敛。

数据增广要克制。手写识别的增广目的是模拟不同扫描设备的差异,而不是制造新的书写风格。我常用的参数范围是:亮度对比度随机乘 0.8 到 1.2,水平平移正负 4 像素,垂直平移正负 2 像素,旋转正负 2 度,透视扰动 1% 以内。笔画粗细的形态学变换只在数据量特别充足时才加,而且概率不超过 0.2。增广后记得存几张图出来,人眼都认不出来的增广样本,模型学出来也是歪的。

3.3 标签编码与 Dataset:CTC blank 与书写者维度划分

文本标签要转成索引序列。CTC 的 blank 索引我固定在 0,<unk>固定在 1。字符表只用训练集构建,全量数据构建字符表会造成标签泄漏,验证集里的低频字在训练时从未出现,却在字符表里占了位置,指标会虚高。

import torch from torch.utils.data import Dataset class HTRDataset(Dataset): def __init__(self, samples, char2idx, target_h=64, max_w=512): # samples: 列表,每个元素是 (img_path, text, writer_id) self.samples = samples self.char2idx = char2idx self.target_h = target_h self.max_w = max_w def __len__(self): return len(self.samples) def __getitem__(self, i): img_path, text, writer_id = self.samples[i] img = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) img = preprocess_line(img, self.target_h, self.max_w) label = [self.char2idx.get(c, self.char2idx['<unk>']) for c in text] return torch.from_numpy(img).unsqueeze(0), torch.tensor(label, dtype=torch.long)

逻辑说明:__getitem__返回的图像张量带上了一个通道维,形状是 (1, 64, W),给后面的 CNN 用。标签不补 padding,交给 DataLoader 的collate_fn去处理,因为同一个 batch 里每个样本的文本长度不一样,统一 padding 需要在 batch 层面做。writer_id在训练循环里不直接用到,但在划分数据集时必须带上。

数据集划分有一个容易忽略的原则:按书写者分,不按行随机分。同一个人的笔迹风格高度一致,如果同一书写者的不同行同时出现在训练集和验证集,验证误差会被严重低估。实际跨书写者评估时,准确率跌 5 到 10 个点都是正常的。划分代码不需要多复杂,但一定得在切分之前就把 writer_id 分组固定好,我一般先把所有样本按 writer_id 聚合,再按 8:1:1 切成训练、验证、测试三份,保证任何一个书写者只出现在其中一份里。

提示:训练集和验证集之间如果发生书写者重叠,模型在验证集上的表现只能证明它记住了这个人的笔迹风格,不能证明它对陌生笔迹的泛化能力。

4. 核心实现与源码解析:CNN + Transformer + CTC 的 PyTorch 最小实现

这一章进入源码部分。整体结构是:CNN 背骨提取视觉特征,把二维特征图压成列序列;位置编码注入水平顺序;Transformer 编码器建模笔画间上下文;最后通过 CTC 头输出字符概率。下面按模块拆开讲。

4.1 CNN 主干改造:保留列特征,丢弃全局语义

ResNet 这类网络原本是为图像分类设计的,最后的全局池化和全连接层拿到的是整张图的语义,但在文本识别里,我们不能把整行压缩成一个向量,而是要保留水平方向每个位置的局部特征。所以常见的做法是砍掉最后一个下采样 stage 和分类头,让输出特征图的高度为 1、宽度保持为原图的八分之一。

import torch.nn as nn from torchvision.models import resnet18 class CNNBackbone(nn.Module): def __init__(self, in_channels=1): super().__init__() self.net = nn.Sequential( nn.Conv2d(in_channels, 64, 7, 2, 3, bias=False), nn.BatchNorm2d(64), nn.ReLU(inplace=True), nn.MaxPool2d(3, 2, 1), ) # 取 resnet18 的 layer1 和 layer2,输出 stride 为 8 res = resnet18(pretrained=True) self.net.add_module('layer1', res.layer1) self.net.add_module('layer2', res.layer2) # 高度压成 1,宽度保持不变 self.pool = nn.AdaptiveAvgPool2d((1, None)) def forward(self, x): # x: (B, 1, H, W) x = self.net(x) # (B, 128, H/8, W/8) x = self.pool(x) # (B, 128, 1, W/8) return x

逻辑说明:第一个卷积层改成了单通道输入,避免 1 通道灰度图在进入 ResNet 之前做无意义的通道复制。layer1和layer2在 torchvision 的 ResNet 定义里分别是两次下采样后的残差块组,输出通道数是 128。AdaptiveAvgPool2d((1, None))里None表示宽度维不限制,这样任意宽度的输入都能进模型,而不会像固定池化核那样强制输出固定宽度。需要留意的是把预训练权重加载到self.net后,layer1 和 layer2 的权重来自 ImageNet,它们对自然图像纹理敏感,但对笔画边缘同样有效;如果数据集足够大,也可以直接从随机初始化训,收敛会慢一些。

4.2 位置编码与 Transformer 编码器:序列化特征的自注意力建模

CNN 输出 (B, 128, 1, W'),先 squeeze 掉高度维,再 permute 成 (B, W', 128),就得到了一个 token 序列。接着用一个线性层把 128 维映射成 d_model 维,再加上位置编码,进入 Transformer 编码器。

import torch import torch.nn as nn class PositionalEncoding(nn.Module): def __init__(self, d_model, max_len=512): super().__init__() self.pe = nn.Embedding(max_len, d_model) def forward(self, x): # x: (B, L, C) B, L, C = x.shape pos = torch.arange(L, device=x.device).unsqueeze(0).expand(B, L) return x + self.pe(pos) class HTRModel(nn.Module): def __init__(self, num_classes, d_model=256, nhead=8, num_layers=4, dim_feedforward=1024, max_len=512): super().__init__() self.backbone = CNNBackbone(in_channels=1) self.proj = nn.Linear(128, d_model) self.pos_enc = PositionalEncoding(d_model, max_len) encoder_layer = nn.TransformerEncoderLayer( d_model, nhead, dim_feedforward, batch_first=True, norm_first=True ) self.encoder = nn.TransformerEncoder(encoder_layer, num_layers) self.ctc_head = nn.Linear(d_model, num_classes) def forward(self, x): # x: (B, 1, H, W) feat = self.backbone(x) # (B, 128, 1, W') feat = feat.squeeze(2) # (B, 128, W') feat = feat.permute(0, 2, 1) # (B, W', 128) feat = self.proj(feat) # (B, W', d_model) feat = self.pos_enc(feat) # 注入水平位置信息 feat = self.encoder(feat) # (B, W', d_model) logits = self.ctc_head(feat) # (B, W', num_classes) return logits

逻辑说明:batch_first=True让 Transformer 编码器直接接受 (batch, seq_len, d_model) 形状,省去手动 permute。norm_first=True是 Pre-Norm 结构,每层先归一化再做注意力,训练更稳定,是当前 Transformer 训练的标准做法。PositionalEncoding用可学习 embedding 而不是正弦编码,原因是手写行的 token 长度集中在几十到一百多,绝对位置编码在这个范围内学起来更灵活。这里把位置编码写成self.pe(pos)加在输入上,没有做 scale,如果训练初期 loss 不降,可以试试把位置编码乘一个 0.1 的系数,减弱初始位置信号对视觉特征的干扰。

4.3 损失计算与推理解码:CTCLoss、贪心去重与 Beam Search

训练时模型的输出要先做log_softmax,转换成对数概率,因为 PyTorch 的 CTCLoss 要求输入是对数概率且形状为 (T, B, C)。这里的 T 应该是有效特征长度,不是右侧 padding 之后的长度。

criterion = nn.CTCLoss(blank=0, zero_infinity=True) def train_step(model, images, labels, seq_lens, optimizer): optimizer.zero_grad() logits = model(images) # (B, T, num_classes) log_probs = logits.log_softmax(-1) log_probs = log_probs.permute(1, 0, 2) # (T, B, num_classes) # T 是模型实际输出的帧数,seq_lens 是每个样本对应的有效帧数 T = log_probs.size(0) input_lengths = torch.full((log_probs.size(1),), T, dtype=torch.long) loss = criterion(log_probs, labels, input_lengths, seq_lens) loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 5.0) optimizer.step() return loss.item()

逻辑说明:input_lengths全部填 T,seq_lens是每个样本的标签长度。这里的关键在于 model 输出时已经把右侧 padding 的帧截断掉,或者用 bucketing 让每个 batch 的 T 取该 batch 内最长的有效宽度,短的样本在 collate 阶段补 pad。CTCLoss 会对超出input_lengths的帧直接忽略,所以即使 logits 里有 padding 帧,只要 input_lengths 设置正确,它们不会参与 loss。不过 Transformer 编码器仍会对 padding 位置做注意力计算,所以最稳妥的办法还是 collate 阶段按 batch 内最窄有效宽截断 logits。

推理阶段最朴素的解码是贪心路径:取每个时间步概率最高的字符,去掉重复,再删掉 blank。实现如下。

def greedy_decode(log_probs): # log_probs: (B, T, num_classes) pred = log_probs.argmax(-1) # 每个时间步取最大概率类别 batch_texts = [] for p in pred: decoded = [] prev = None for idx in p.tolist(): if idx != 0 and idx != prev: # 0 是 blank,连续重复只保留一个 decoded.append(idx) prev = idx batch_texts.append(decoded) return batch_texts

逻辑说明:CTC 的去重规则是“相同字符连续出现只算一次”,这个特性对中文手写是友好的,因为正常文本里很少出现连续两个完全相同的汉字;如果语料里有叠字标注,就需要在标签里插入 blank 分隔符,或者在解码后处理里做特殊判断。贪心解码速度快、实现简单,作为 baseline 够用;想要更高准确率,可以用 beam search 在每一步保留概率最高的若干条路径,合并相同前缀,等验证指标稳定后再加。torchaudio 或其他语音识别工具库里有现成的 CTC beam search 实现,直接调用比自己维护一份前缀合并逻辑省事。

5. 避坑与排查:手写识别翻车现场的 4 个高频根因

这个领域翻车点很集中,练到后面你会发现翻来覆去就是那几个原因。下面四条是我实际调模型时常碰到的,每条按现象、原因、解决三个步骤写。

5.1 多余 padding 帧给 CTC 送分:重复输出与整行空白

现象:训练 loss 降得很快,但模型预测结果要么是整行空白,要么是末尾字符反复出现。

原因:collate 阶段把不同宽度的样本统一 pad 到 batch 最大宽度,padding 帧也被送进 Transformer 编码。CTC 看到 padding 位置全是白底特征,输出大量 blank,模型学会把“右边没有内容”当成“输出空字符”,整行解码结果被空白吞掉。另一个常见变体是input_lengths误填成 batch 最大宽度,而没有减去右侧 padding 的帧数,真实有效的尾部帧被白白忽略。

解决:把宽 512 上限收紧到 384,超宽样本在数据预处理时就截断;collate 时记录每个样本的真实有效宽度,在下采样后得到有效帧数,input_lengths用它来设。同时把max_w相关的 padding 放在右侧,不要左右均衡填充,这样 padding 集中在序列末尾,和 CTC 忽略尾部帧的逻辑正好对齐。我是把preprocess_line的返回值顺带多返回一个有效宽度字段,而不是在 Dataset 里重新算。

5.2 没有 warmup 的 Transformer 是玄学:损失冲高与不收敛

现象:前 200 步 loss 在十几和几十之间剧烈震荡,然后某一步直接变 NaN;或者 loss 缓慢下降但永远到不了正常水平。

原因:Transformer 的参数量集中在多头注意力和 FFN 上,Adam 在训练初期的一阶矩估计还没站稳,较大的初始学习率会直接把权重推到数值不稳定的区域。CNN + BLSTM + CTC 时代常用的 1e-3 学习率在 Transformer 上完全不适用,这不是模型写错了,而是学习率策略没跟上。

解决:用 AdamW,峰值学习率控制在 1e-4 到 3e-4 之间,前面加 500 步左右的 warmup。warmup 阶段学习率从 0 线性升到峰值,后面接 cosine 衰减。

import math def build_lr_scheduler(optimizer, warmup_steps, total_steps): def lr_lambda(step): if step < warmup_steps: return step / warmup_steps t = (step - warmup_steps) / max(1, total_steps - warmup_steps) return 0.5 * (1 + math.cos(math.pi * t)) return torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)

逻辑说明:total_steps根据数据集大小、batch size 和训练轮数提前算好。如果 batch size 很大,峰值学习率可以拉到 3e-4;batch size 只有 8 或 16 时,1e-4 更稳。warmup 步数不是越多越好,小数据集 300 步够用,大规模数据 1000 步也不嫌多。

5.3 长文本行撑爆显存:自注意力的平方开销

现象:训练时突然报 CUDA OOM,定位后发现是某个 batch 里有一张特别宽的行图;把 batch size 调小后,短样本的 GPU 利用率又上不来。

原因:Transformer 编码器的注意力复杂度是序列长度的平方。图像宽度 768px 时,下采样 8 倍后只剩 96 个 token,还没问题;一旦出现宽度 1536px 的行图,token 数到 192,注意力矩阵从 9216 个元素涨到 36864 个元素,显存开销翻了四倍,再碰上大 batch 直接爆。

解决:两招一起用。第一,数据预处理里把max_w限制在 512,更长的文本行按比例缩放或截断;第二,训练时用宽度分桶,把宽度接近的样本放进同一个 batch,batch 内的最大宽度只比最小宽度高一小截,这样每个 batch 的序列长度不会出现极端波动。分桶方式不复杂:离线阶段先按图像宽度排序,然后按 batch 大小切成若干块,块内随机打乱,collate 时取块内最大宽度作为统一宽度,其余补 pad。序列长度波动被限制住之后,显存占用立刻变得可控。

5.4 增广失控让笔画断裂:训练 loss 好、验证字准率差

现象:训练集准确率很快到 95%,验证集却卡在 80% 上不去,抽 badcase 一看,好多字形结构明显不完整,像字被切掉一笔。

原因:增广参数设得太激进。旋转超过 5 度会让竖直笔画变成斜线,形态学腐蚀会让细笔画断开,随机擦除类增强会直接删掉关键笔画。手写识别和图像分类不同,分类删除部分像素可能不影响类别判断,但汉字删掉一个横或一个撇就是另一个字或者破字,模型被迫去学“残缺字形也能猜”,泛化能力反而不升。

解决:增广要模拟扫描设备差异,不是模拟书写者。旋转控制在正负 2 度以内,透视扰动 1%,亮度对比度乘 0.8 到 1.2,平移不要超过 4 像素,不做随机擦除、不做大面积遮挡。每轮增广后保存一批可视化结果,把原图和增广图并排看一眼,人眼都能轻松辨认,再放心拿去训练。这个检查成本很低,但能省去后面大量猜 badcase 的时间。

6. 验证与进阶:用编辑距离测准模型,再谈部署

手写识别最常用的验证指标是字符错误率 CER 和整行准确率。CER 用编辑距离除以真实文本长度,比单纯看准确率更能反映模型在长文本行上的退化程度。下面这段代码不依赖外部库,直接算编辑距离:

def cer(gt, pred): # gt 和 pred 都是字符串 dp = [[0] * (len(pred) + 1) for _ in range(len(gt) + 1)] for i in range(len(gt) + 1): dp[i][0] = i for j in range(len(pred) + 1): dp[0][j] = j for i in range(1, len(gt) + 1): for j in range(1, len(pred) + 1): cost = 0 if gt[i - 1] == pred[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[-1][-1] / len(gt)

评估时先按书写者分组算 CER,再求整体均值,不要让同一份测试集里出现同一个人的多行,否则数字会好看很多。进阶的调试习惯是按 CER 从高到低把 badcase 排出来,看错误集中在哪种字上。我常遇到的情况是:错误集中在笔画数多的复杂字和手写潦草的行,前者可以补充对应字体的训练样本,后者说明图像分辨率不够,需要把输入高度从 64 提到 128。改模型结构之前先看 badcase,几乎每次都比换解码器更划算。

部署阶段有一个细节值得提前想清楚:Transformer 编码器在 ONNX 导出时序列维是动态的,导出代码里必须把sequence轴声明为动态轴,否则模型只能接受固定宽度的输入,生产环境一旦出现更宽的图就得切成多段。半精度推理时先验证一下位置编码层的数值范围,fp16 下位置 embedding 的精度损失有时会在长序列上放大,导致输出出现连续重复字符。我现在的习惯是:先用 5000 行小数据跑通整个流程,看到模型能过拟合,再上全量数据;任何结构改动都把训练曲线和 badcase 截图存档,对比时心里有数。手写识别没有银弹,把数据管线和验证工具做扎实,比反复换模型结构提升来得更实在。希望帮到你。

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

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

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

立即咨询