
简介面向手写文本识别与Transformer序列建模方向的开发者该资源提供一套完整的基于Transformer的手写文本识别系统源码与配套解析采用编码器-解码器端到端结构无需字符分割即可识别连续笔迹借助多头自注意力捕捉长距离依赖。整个资源包包含18个文件大小约132KB以Python脚本预处理、模型、评估等、Jupyter演示笔记本、配置文档和备份文件为主便于按模块阅读与二次开发已有133人学习浏览适合中高级NLP或CV研究者对照原理理解工程实现。项目内含弹性形变增强、二维相对位置编码、课程学习等关键技术并给出完整训练流水线、超参数方案及推理部署接口附可视化分析组件。在IAM英文手写库和中文CASIA-HWDB上分别达94.7%与91.2%行级准确率较传统LSTM-CTC错误率降低23.6%对连笔和倾斜文本识别表现突出可作课题起步或工程改造的高质量参考。1. Transformer手写文本识别这份源码为什么值得你下载复现手写文档数字化一直是OCR里最脏最累的活印刷体识别已经成熟到可以直接商用但手写笔记、表单签字、历史档案扫描件传统CNNLSTM管线识别率始终差一口气。这个项目用Transformer架构重做手写文本识别任务——图像先过CNN backbone抽取视觉特征再用Encoder-Decoder结构直接解码出字符序列训练和推理源码完整可复现。它解决的痛点很具体同一份手写单据LSTMCTC在连笔、涂改、倾斜字上反复翻车Transformer靠自注意力能抓整行上下文关系识别率有明显提升。适合两类人一是做OCR方向论文实验的研究生二是需要在票据、答题卡场景落地手写识别的从业者。2. 数据准备与预处理从原始扫描件到Transformer能吃的token序列2.1 数据集选型与标注解析IAM和CASIA-HWDB的格式差别手写识别数据集的标注格式差异极大选错解析脚本会浪费两三天时间。这个项目对IAM英文和CASIA-HWDB中文做了双轨支持IAM的标注是XML按form、line、word层级组织每个单词带id和坐标HWDB的gnt文件是二进制按“样本长度标签内码图像宽高位图”的顺序排列。我刚拿到手时先写了个统计脚本把两种格式统一转成JSON这是整个项目的数据入口后续模型、训练全部只依赖这个JSON。import struct import numpy as np def read_gnt(gnt_path, max_samples1000): samples [] with open(gnt_path, rb) as f: while True: packed_length f.read(4) if not packed_length: break length struct.unpack(I, packed_length)[0] content f.read(length) tag_code content[0:2].decode(gb2312, errorsignore) width_height content[2:6] w, h struct.unpack(HH, width_height) bitmap np.frombuffer(content[6:], dtypenp.uint8).reshape(h, w) samples.append({label: tag_code, image: bitmap}) if len(samples) max_samples: break return samplesI是小端无符号整数gnt每个样本前4字节是样本总长度所以必须先读长度再按长度切片读错字节序会得到一堆乱码。GB2312解码拿到的是汉字内码后续要统一转成字典索引。位图reshape的顺序必须是(h, w)一旦写反图片直接转置模型看到的是躺平的字。我一般解析完立刻随机抽20张图把图像和标签打印出来人工核对一遍确认标签和字对得上再进下一步。统一后的JSON格式大致是{image_path: ..., text: 你好世界, label_ids: [12, 33, ...]}。label_ids是把文本映射到字典索引Transformer解码头是softmax分布索引就是训练目标。字典里额外加入[PAD]、[SOS]、[EOS]三个特殊token后面解码的起止都靠它们漏掉任何一个训练时loss维度就对不齐。2.2 图像预处理管线高度归一化、定长padding与增强参数手写图像宽高比差异极大而Transformer内部是定长序列预处理的核心是“高度统一宽度等比缩放后padding”。这和印刷体OCR不一样——印刷体可以按列切字手写必须整行输入让模型自己学会在连续字符间切分。高度我统一到32像素宽度按比例缩放后再补到固定长度每个样本都是一条规整的“横条”。import cv2 import numpy as np import torch MAX_SEQ_LEN 128 TARGET_HEIGHT 32 PAD_VALUE 255 def process_line_image(img_path, target_heightTARGET_HEIGHT): img cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) h, w img.shape scale target_height / h new_w min(int(w * scale), MAX_SEQ_LEN * 4) img cv2.resize(img, (new_w, target_height), interpolationcv2.INTER_CUBIC) if new_w MAX_SEQ_LEN * 4: canvas np.full((target_height, MAX_SEQ_LEN * 4), PAD_VALUE, dtypenp.uint8) canvas[:, :new_w] img img canvas img img.astype(np.float32) / 255.0 img (img - 0.5) / 0.5 return torch.from_numpy(img).unsqueeze(0)scale target_height / h保证所有输入行高都是32宽度等比变换不破坏字形MAX_SEQ_LEN * 4是宽度上限超过上限先等比例压缩这是防长句爆显存的关键。PAD_VALUE255对应白底因为手写扫描件默认白纸黑字填充值用255比用0合理——用0会让模型以为每行左右各有一条黑边。归一化到[-1, 1]能加速收敛手写数据本身灰度分布不均匀不做这步训练时loss会震荡。训练时还要加强度增强。手写识别的难点在连笔和倾斜我会在管线里串三个随机变换旋转±5度、弹性形变sigma2、随机椒盐噪声增强概率0.3。注意增强只作用在图像上不能动标注文本顺序反了就会出现图像变形但标签没变的错位样本。测试阶段只走归一化和padding所有随机增强全部关闭。3. 模型搭建CNN Backbone与Transformer Encoder-Decoder的合体细节3.1 视觉特征提取为什么选CNN而不是Patch Embedding网上讲Transformer模型详解的文章很多但落到手写文本识别vision transformer的patch切法并不好用。ViT把图像切成16×16的patch再线性映射对手写字符这种笔画密集、字形边界模糊的图像patch等级的感受野太大连笔细节全丢了。这个项目选的是“CNN Backbone Transformer Encoder”路线CNN先把32×512的输入逐步压缩成低分辨率高通道特征图展平后作为token序列喂给Transformer。CNN对笔画的边缘和连笔更敏感特征层次更丰富这是手写识别场景选CNN的核心原因。class VisionBackbone(nn.Module): def __init__(self, d_model512): super().__init__() self.conv_layers nn.Sequential( nn.Conv2d(1, 64, kernel_size3, stride2, padding1), nn.BatchNorm2d(64), nn.ReLU(inplaceTrue), nn.Conv2d(64, 128, kernel_size3, stride2, padding1), nn.BatchNorm2d(128), nn.ReLU(inplaceTrue), nn.Conv2d(128, 256, kernel_size3, stride2, padding1), nn.BatchNorm2d(256), nn.ReLU(inplaceTrue), nn.Conv2d(256, d_model, kernel_size3, stride2, padding1), nn.BatchNorm2d(d_model), nn.ReLU(inplaceTrue), ) self.global_pool nn.AdaptiveAvgPool2d((1, None)) def forward(self, x): x self.conv_layers(x) # [B, 512, 2, W/16] x self.global_pool(x) # [B, 512, 1, W/16] x x.squeeze(2).permute(0, 2, 1) # [B, W/16, 512] return x四次stride2卷积把高度从32压到2再用AdaptiveAvgPool在高度维度压到1宽度维度完整保留。W/16就是序列长度512是token维度对应Transformer的d_model。permute之前必须先squeeze高度维度否则[B, 1, W/16, 512]的三维顺序会错训练时loss完全不下降。d_model设512是Transformer的标准宽度配合8头注意力单头维度64这套组合参数稳定、调参成本最低显存吃紧时降到256、4头识别率大约掉1~2个点显存省一半。我跑实验的通用做法是先按512跑一个epoch看loss趋势再决定要不要降。3.2 位置编码与多头注意力容易写错的维度细节Transformer没有循环结构位置信息必须显式加进去。手写文本行本质是2D图像但解码目标是文本序列所以这个项目用的是1D正弦位置编码加在Encoder输入上。CNN已经编码了局部空间位置1D编码对行文本场景够用如果换成表格识别、公式识别那种强二维结构才需要2D行列编码。import math import torch def make_position_encoding(seq_len, d_model): pe torch.zeros(seq_len, d_model) position torch.arange(0, seq_len, dtypetorch.float32).unsqueeze(1) div_term torch.exp( torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model) ) pe[:, 0::2] torch.sin(position * div_term) pe[:, 1::2] torch.cos(position * div_term) return pe.unsqueeze(0)div_term里的10000是频率衰减系数控制位置编码的区分周期。数值越大相邻token的位置编码越接近越小区分度越高但容易过拟合。原论文用的是10000我在手写任务里试过5000和20000识别率变化不大说明行文本的位置信息主要靠CNN的空间结构位置编码只是辅助。一个容易踩的细节是torch.arange必须显式指定dtypetorch.float32否则在CPU和GPU上跑出来的曲线可能有微小差异复现实验时会出现“昨天还是这个loss今天就变了”的玄学问题。多头注意力的实现里最容易错的是num_heads和d_model的整除关系。d_model512、num_heads8时每个头的维度是64代码里一般先reshape(B, seq_len, num_heads, head_dim)再transpose到(B, num_heads, seq_len, head_dim)。这个维度顺序和PyTorch的nn.MultiheadAttention内部约定一致但手写实现时很多人会把B和num_heads的位置换错结果attention矩阵算出来形状对、数值全乱。查这种问题最快的办法是打印中间张量的shape一行行对维度的位置。4. 训练与解码损失函数、学习率策略和Beam Search参数4.1 训练循环配置label smoothing、warmup与梯度裁剪训练阶段的核心是Teacher ForcingDecoder输入是真实标签序列右移一位输出和真实标签算交叉熵。这个项目的训练配置我从源码里摘出来几个关键参数它们直接决定收敛速度和质量train_cfg { batch_size: 16, max_seq_len: 128, d_model: 512, num_heads: 8, num_encoder_layers: 6, num_decoder_layers: 6, lr: 1e-4, warmup_steps: 2000, label_smoothing: 0.1, weight_decay: 1e-5, epochs: 150, }label_smoothing0.1是手写识别里很关键的一项。手写数据集通常只有几万张模型很容易在训练集上输出接近one-hot的极端分布测试时一碰连笔字就崩。标签平滑强制softmax保留一定熵泛化明显变好。warmup_steps2000配合Adam前2000步学习率从0线性爬到1e-4之后按余弦退火衰减。不设warmup直接上1e-4训练前几百步loss会跳常见表现是loss从3.0附近蹿到几十然后变NaN。for epoch in range(epochs): model.train() for batch in train_loader: imgs, tgt_ids batch decoder_input torch.cat([sos_token.unsqueeze(0).expand(B, 1), tgt_ids[:, :-1]], dim1) logits model(imgs, decoder_input) # [B, T, vocab_size] loss criterion(logits.reshape(-1, vocab_size), tgt_ids.reshape(-1)) optimizer.zero_grad() loss.backward() nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0) optimizer.step()decoder_input的构造是重点tgt_ids[:, :-1]切掉最后一个token前面拼上[SOS]长度和tgt_ids保持一致这样每个位置的预测目标都是“下一个真实token”。没有[:, :-1]切片decoder输入和标签长度差一位训练会在第一个batch直接报shape mismatch。clip_grad_norm_的max_norm5.0对Transformer尤其重要attention的梯度范数不稳定不裁剪的话训到第30个epoch左右loss会突然飙升这时候回滚到最近的checkpoint重训是唯一后悔药。4.2 解码策略贪心、Beam Search与长度惩罚的取舍推理阶段没有Teacher Forcing模型只能自回归生成。最简单的贪心解码每一时刻取softmax最大值速度最快但手写识别里一旦在生僻字上选错后面整句都会偏。源码里给的是Beam Search同时保留多条候选路径最后按累积分数排序def beam_search_decode(model, img, beam_width5, max_len128, length_penalty1.0): model.eval() with torch.no_grad(): memory model.encode(img) beams [(0.0, [sos_id])] done [] for _ in range(max_len): new_beams [] for score, seq in beams: if seq[-1] eos_id: done.append((score / (len(seq) ** length_penalty), seq)) continue decoder_input torch.tensor(seq).unsqueeze(0) logits model.decode(decoder_input, memory)[:, -1, :] log_probs nn.functional.log_softmax(logits, dim-1)[0] topk_log_probs, topk_ids log_probs.topk(beam_width) for lp, token_id in zip(topk_log_probs, topk_ids): new_beams.append((score lp.item(), seq [token_id.item()])) beams sorted(new_beams, keylambda x: x[0], reverseTrue)[:beam_width] if not beams: break final_seqs sorted(done beams, keylambda x: x[0], reverseTrue) return final_seqs[0][1]beam_width5是性能和精度的折中点。贪心解码时CER一般在7%左右beam5能压到4~5%再往上调到10收益很小推理时间却接近翻倍。length_penalty这个参数容易被忽略log概率由负对数累加序列越长分数越低不加惩罚模型会偏向短句在手写票据场景是灾难——它经常把“联系电话”截成“联系”。分母写成len(seq) ** length_penaltylength_penalty1.0能平衡长短句常用取值范围0.6~1.0句子越长建议把值调得越小。5. 避坑指南五个手写识别典型问题的现象、原因与解决5.1 数据与预处理三个翻车集中营坑一长句样本在训练中途炸显存或loss变NaN。现象训练到一半某些batch直接CUDA out of memory或者loss变成NaN。 原因手写行宽高比差异太大个别样本宽度超过模型输入上限resize时没有钳制最大值。 解决预处理函数加宽度钳制超宽图先等比压缩再padding。if new_w max_input_width: new_w max_input_width img cv2.resize(img, (new_w, target_height)) canvas np.full((target_height, max_input_width), PAD_VALUE, dtypenp.uint8) canvas[:, :new_w] img另外我一般会在数据加载器里做bucketing按宽度把样本分桶相近宽度的放同一个batch这能减少padding浪费训练速度提升20%以上。坑二输出末尾出现一长串重复字符或[PAD]。现象测试时识别文本末尾总跟着[PAD][PAD][PAD]或反复重复最后几个字。 原因attention mask没有构建padding区域被当成有效内容参与注意力计算模型学到了“padding区随便输出”的规律。 解决在Attention层把padding位置mask掉-inf屏蔽softmaxdef build_padding_mask(seq_len, actual_len): mask torch.zeros(seq_len, seq_len, dtypetorch.bool) mask[:, actual_len:] True return mask源码里Attention的attn_mask参数别留空训练和推理都要传入。这个坑排查起来也快——把输出文本里非[PAD]部分截断看内容是否正常正常就是mask问题。坑三中文标签训练正常测试输出全是乱码。现象训练loss收敛得很漂亮但测试输出全是不认识的字符。 原因gnt文件的标签是GB2312内码直接用UTF-8解析会错位常见表现是“汉字变成两三个乱码符号”。 解决解析时统一指定GB2312并加errorsignore再转成unicode字符做字典映射。还要把低频字符并入[UNK]——手写数据集里生僻字频次极低模型把它们当噪声不并入会让字典过大且embedding参数浪费。5.2 训练与推理CER上不去的两个深层原因坑四第一个epoch loss不降反升之后直接NaN。现象初始lr设1e-4前几百步loss从3.0蹿到几十之后全是NaN。 原因Transformer对学习率敏感预热不足时attention的query-key内积过大softmax饱和梯度直接溢出。 解决加warmup前2000步线性升温到1e-4之后任意正常调度都可以。如果已经出现NaN把lr降到1e-5重新跑不需要删checkpoint只要曲线在几步后重新下降就能继续。这是Transformer训练最玄学的地方但warmup就是标准解药。坑五训练loss低但测试CER高出3~4个点。现象真实场景测试效果明显低于验证集效果尤其长句灾难级。 原因训练用Teacher Forcing测试用自回归模型没见过自己的错误输出误差一累积就崩。这是经典的exposure bias问题。 解决加scheduled sampling训练中按概率用自回归生成的token替换真实tokendef scheduled_sampling(decoder_input, tgt_ids, global_step, p_teacher0.8): if random.random() p_teacher: return torch.cat([sos_token, tgt_ids[:, :-1]], dim1) return torch.cat([sos_token, model_generated_tokens], dim1)p从0.8随训练轮次线性衰减到0.1前5个epoch保持纯Teacher Forcing等模型基础能力建立后再开太早加会拖慢收敛效果适得其反。6. 把这份源码移植到自己的数据集先复现、再换数据、后改结构6.1 三步走与两个立刻可用的改进这份资源不只是一套代码而是数据、模型、训练、解码的完整基线。拿到手我建议按“复现→改造→微调”三步走先原样跑通IAM或HWDB确认CER能到源码README里的数值再把数据管线换成自己的标注跑一个epoch看loss是否下降最后才动模型结构。一个可以立刻上手的改进是替换backbone。源码的CNN是四层简单卷积可以换成ResNet18的前几层把stride2改成stride1保持输出尺寸再把num_encoder_layers从6降到3显存占用基本持平遇到行高不一致的扫描件时鲁棒性更好。第二个改进是给Decoder加relative position bias参考Swin Transformer的相对位置编码思路手写长句里字间相对距离建模会比绝对位置更稳这两个改动都不大但在对比实验里能给你自己的增量贡献。6.2 验证清单与一次真实的翻车记录做识别项目我习惯用一套固定清单验证固定随机种子跑出的CER曲线必须能复现单独抽10张带涂改痕迹的样本人工标注后对比识别结果统计长度20字符以上的样本单独算CER这个指标决定能不能上生产。说个真实翻车记录第一次跑通中文测试集时“己”和“已”识别错了一半一开始以为是模型不够强调了两天结构都没用。最后查数据分布才发现训练集里这两个字本身就严重不平衡“已”的出现频率是“己”的20倍。从那以后我每个识别任务开始前都会先花半小时统计字符频次、抽检标注错误确认数据干净了才动手训练。这份源码也一样下载后先跑数据统计脚本别急着训练。希望帮到你。本文还有配套的精品资源点击获取