ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

CRNN文字识别模型:CNN+BiLSTM+CTC原理与PyTorch实现解析

2026/9/16 21:18:26 拓冰建站 浏览量
CRNN文字识别模型:CNN+BiLSTM+CTC原理与PyTorch实现解析 简介一份CRNN文字识别完整PyTorch实现面向深度学习与OCR场景文字识别方向的开发者、学生重点解决端到端文字识别训练和不定长文本序列预测问题融合CNN与RNN结构无需文字预先分割。资源包共2000个文件解压后约107.78MB以大量png图片样本、py训练/预测脚本、mat数据文件、ipynb演示为核心同时附带说明文档与模型文件结构清晰。作者基于IIIT-5k数据集完成训练模型中已覆盖训练与预测流程可直接调用ipynb部分还展示了利用PyTorch搭建CRNN进行验证码识别支持自定义图像输入与网络结构调整可灵活用于实验拓展。已有2582人学习过该资源适合希望从原理到实战完整掌握CRNN文本识别的读者。1. 为什么说 CRNN 是文字识别绕不开的基线模型做过 OCR 工程的人都知道2015 年提出的 CRNN 到今天依然是一个绕不开的模型它把 CNN 的特征提取能力和 RNN 的序列建模能力拼在一起用 CTC 损失做端到端训练输入一张图片直接给出字符串不需要预先切分字符也不需要把检测和识别拆成两套系统。这份资源里带着完整源码、IIIT-5k 训练数据和训练好的权重就连验证码识别场景也给出了一个跑通的可视化 notebook。如果你是刚接触深度学习文字识别的新手可以用它把论文里的结构一条条对到代码上如果你已经在做 OCR 落地里面定宽裁剪、字符边界处理、CTC 解码这些代码仍然值得翻一翻。它解决的核心问题很具体任意长度文本的识别如何在不需要逐字标注的情况下完成训练和推理。开源中文 OCR 领域里 CRNN 相关的实现很多但这份附带的 ipynb 把 PyTorch 训练过程完整串了一遍适合直接在此基础上改。2. CRNN 网络骨架CNN 特征提取、BiLSTM 序列建模与 CTC 对齐2.1 为什么是 CNN RNN而不是纯 CNN场景文字识别的难点在于字符宽度不固定、图像长度不定。纯 CNN 做分类需要先把图片裁剪成固定尺寸对长文本就无能为力了。CRNN 的思路是把 CNN 当作特征提取器输出一个高度压缩、宽度保留的特征序列然后交给双向 LSTM 去建模字符之间的上下文依赖最后用 CTC 解决序列长度对不上的问题。严格说CNN 部分不是拿来直接分类的而是把图像转换成按时间步排列的特征向量序列。数据集里出现的traindata.mat、testCharBound.mat这些文件本质上也是围绕这个设计组织的图片数据和字符边界数据分开存储训练时既要知道图上有什么字也要知道字大概在什么位置。模型本身并不依赖边界做训练但边界信息可以用来验证对齐效果和生成可视化结果。2.2 从原始图像到特征序列的关键变换CRNN 输入图像高度固定为 32宽度可以任意。经过卷积和池化后特征图的宽度大约是原始宽度的 1/4这个 1/4 很关键因为 CTC 要求输入序列长度和输出标签长度有一个可学习的对应关系但序列每一帧仍然覆盖多个原始像素宽度。下面是 PyTorch 实现的核心网络结构与论文中的配置保持一致import torch import torch.nn as nn class CRNN(nn.Module): def __init__(self, n_classes, hidden_size256): super().__init__() # CNN 部分把单通道灰度图映射成高层视觉特征 self.cnn nn.Sequential( nn.Conv2d(1, 64, 3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2, 2), # 高度、宽度各减半 nn.Conv2d(64, 128, 3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d(2, 2), # 高度、宽度再减半 nn.Conv2d(128, 256, 3, padding1), nn.BatchNorm2d(256), nn.ReLU(inplaceTrue), nn.Conv2d(256, 256, 3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d((2, 1)), # 只压缩高度保留宽度 nn.Conv2d(256, 512, 3, padding1), nn.BatchNorm2d(512), nn.ReLU(inplaceTrue), nn.Conv2d(512, 512, 3, padding1), nn.ReLU(inplaceTrue), nn.MaxPool2d((2, 1)), # 只压缩高度 nn.Conv2d(512, 512, 3, padding1), nn.BatchNorm2d(512), nn.ReLU(inplaceTrue), ) # RNN 部分双向 LSTM 建模序列上下文 self.lstm nn.LSTM(512, hidden_size, num_layers2, bidirectionalTrue, batch_firstTrue) self.fc nn.Linear(hidden_size * 2, n_classes) def forward(self, x): feat self.cnn(x) # [B, 512, h, w] B, C, H, W feat.shape seq feat.reshape(B, C * H, W) # 把高度通道合并 seq seq.permute(0, 2, 1) # [B, W, 512] out, _ self.lstm(seq) # [B, W, 512] return self.fc(out) # [B, W, n_classes]这里把最后一层特征图reshape成C * H维向量目的是兼容输入图像高度不是严格 32 的情况。通常输入高 32、宽 W 的图像经过三次高度方向的池化后特征图高度会变成 2所以reshape后序列长度是W / 4每个时刻的特征维度是512 * 2 1024。双向 LSTM 的隐藏层维度设 256双向拼接后恰好也是 512全连接层直接映射到字符类别数n_classes。注意n_classes包含一个 CTC 的 blank 类别所以实际字符表外要加 1。2.3 CTC 损失解决的不只是长度对齐模型输出的序列长度与真实标签长度不一致这是训练阶段的根本矛盾。CTC 的做法是维护一个带 blank 的路径集合blank 表示当前位置没有字符通过前向-后向算法把所有合法对齐路径的概率求和再取负对数作为损失。代码库中的ctc_pytorch_tensorboard.ipynb正是用 PyTorch 内置的CTCLoss在做这件事。下面给出训练循环里的核心调用方式criterion nn.CTCLoss(blank0, zero_infinityTrue) # 假设 batch 内只有一张图output 形状 [B, T, C] output model(img) # [1, T, n_classes] T output.size(1) loss criterion( output.log_softmax(2).transpose(0, 1), # 需要 [T, B, C] targets, # 拼接后的标签索引 input_lengths, # 每个样本的序列长度 target_lengths # 每个样本的标签长度 )blank0表示类别索引 0 被保留给空字符这也意味着字符表构建时要从 1 开始编号。zero_infinityTrue是个容易被忽略的细节当某个 batch 内输入序列长度小于目标长度时CTC Loss 会算出正无穷这个开关把它置为 0避免整个训练直接发散。input_lengths在定宽裁剪场景下通常是W / 4但如果做了动态 batch就必须逐样本计算。代码里把output.log_softmax(2)放在transpose之前是因为三维概率已经被网络输出过了只需对类别维做 softmax。下表总结了输入宽 W 时各关键张量的尺寸变化排查维度不匹配时直接对照它位置张量形状含义输入图像[B, 1, 32, W]固定高度 32宽度可变CNN 输出[B, 512, 2, W/4]高度压缩到 2序列输入 LSTM[B, W/4, 1024]高度通道合并LSTM 输出[B, W/4, 512]双向拼接分类输出[B, W/4, n_classes]每个时刻一个类别分布3. 数据管线与训练从 IIIT-5k 到 train_fix_width.pkl3.1 源码里这些文件分别承担什么角色第一次打开这个项目时先别急着跑训练把数据文件之间的关系理清楚能省很多调试时间。traindata.mat和testdata.mat存的是 IIIIT-5k 数据集的合成图片矩阵trainCharBound.mat和testCharBound.mat存的是每个字符在图片中的边界框列表。训练时真正喂给模型的其实是train_fix_width.pkl它把图片统一处理成了固定宽度并和标签索引一一对应。文件和用途对照文件内容用途traindata.mat训练图像矩阵原始数据需要转成 png 或 npytestdata.mat测试图像矩阵评估模型用trainCharBound.mat训练集字符边界验证对齐、生成可视化testCharBound.mat测试集字符边界评估边界还原精度train_fix_width.pkl固定宽度处理后的训练样本直接作为训练集输入ctc_pytorch_tensorboard.ipynb完整训练 TensorBoard 可视化验证码识别实验3.2 定宽处理为什么不能直接 resize有些新手会把所有图片直接缩放成一个固定宽高比比如32 x 280。这种做法对 CRNN 是有害的字符本身的长宽比被破坏模型学到的字符特征会在推理时失真。正确做法是先保持高宽比缩放到高度 32然后对不足固定宽度的部分做 padding超过的部分做适度压缩。源码里的train_fix_width.pkl应该就是在这一逻辑下生成的。下面是一个具备同样行为的 Dataset 实现片段import cv2 import torch from torch.utils.data import Dataset class OCRDataset(Dataset): def __init__(self, samples, char_dict, img_height32, fix_width280): self.samples samples # [(img_array, label_str), ...] self.char_dict char_dict # 字符到索引的映射0 留给 blank self.img_height img_height self.fix_width fix_width def __len__(self): return len(self.samples) def __getitem__(self, idx): img, label self.samples[idx] if not isinstance(img, torch.Tensor): img torch.from_numpy(img).float() h, w img.shape scale self.img_height / h new_w int(w * scale) img img.unsqueeze(0).unsqueeze(0) # [1, 1, H, W] img torch.nn.functional.interpolate( img, size(self.img_height, max(new_w, 1)), modebilinear ).squeeze(0) # [1, 32, new_w] # 宽度不足时右侧补零 if new_w self.fix_width: pad torch.zeros(1, self.img_height, self.fix_width - new_w) img torch.cat([img, pad], dim2) else: img img[:, :, :self.fix_width] target torch.tensor([self.char_dict[c] for c in label], dtypetorch.long) return img, target这里的interpolate是等比缩放到高度 32 的关键fix_width一般取数据集中最长样本的宽度过大会浪费算力过小会被截断。右侧补零是常见做法但要注意 padding 区域在训练早期容易让模型学到右边永远是空白的偏置因此不少实现会在 padding 区域随机填充噪声。源码里2332_2.png这类样本可以直接读进来作为调试数据验证预览时看到的和模型输入是否一致。3.3 变长 batch 的 collate_fn 怎么设计CRNN 的 batch 内图片宽度不同不能简单用默认的collate_fn堆叠。常见做法是先把 batch 内所有样本按宽度从大到小排序然后取最大宽度做 pad。排序的作用是让 CTC 的input_lengths计算更直观也方便在推理阶段做按需裁剪。配套的collate_fn如下def collate_ocr(batch): imgs, targets zip(*batch) max_w max(img.shape[2] for img in imgs) img_tensor torch.zeros(len(imgs), 1, 32, max_w) for i, img in enumerate(imgs): img_tensor[i, :, :, :img.shape[2]] img target_concat torch.cat(targets) target_lens torch.tensor([len(t) for t in targets], dtypetorch.long) return img_tensor, target_concat, target_lenstargets被拼接成一个一维张量target_concat这是因为CTCLoss接受扁平化的标签序列配合target_lengths才能把每个样本的边界切出来。这里没有显式传入input_lengths是因为当前 batch 已统一 pad 到max_w序列长度都是max_w // 4。这种做法在 batch 内部宽度差距很大时会浪费计算资源工程上更激进的做法是直接按宽度分桶。3.4 训练脚本里必须调好的几个参数模型和数据都就绪后训练环节最容易犯的错误集中在三个地方学习率、CTC 的 blank 索引、以及序列长度的下界。IIIT-5k 这类合成数据相对干净Adam 优化器配1e-3初始学习率通常能正常收敛但迁移到自己采集的数据时建议降到1e-4。from torch.utils.data import DataLoader model CRNN(n_classeslen(char_dict) 1) optimizer torch.optim.Adam(model.parameters(), lr1e-3) scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size10, gamma0.1) loader DataLoader(dataset, batch_size32, shuffleTrue, collate_fncollate_ocr, drop_lastTrue) for epoch in range(30): for imgs, target_concat, target_lens in loader: output model(imgs) # [B, T, C] T output.size(1) input_lens torch.full((imgs.size(0),), T, dtypetorch.long) loss criterion(output.log_softmax(2).transpose(0, 1), target_concat, input_lens, target_lens) optimizer.zero_grad() loss.backward() optimizer.step() scheduler.step()step_size10配合gamma0.1是 OCR 任务里常见的阶梯式下降策略epoch 10 之前模型还在学习字符的局部纹理过早降学习率会让 RNN 部分难以学到长距离依赖。drop_lastTrue是为了防止最后一个 batch 样本数过少导致input_lengths和target_lengths出现极端值。训练过程中如果发现 loss 一直在 1 附近震荡优先检查字符表里是否混入了重复字符或全角符号这类噪声会让 CTC 的概率分布长期无法收敛到单峰。4. 推理解码贪心搜索、束搜索与字符边界还原4.1 解码的本质是从概率分布中恢复文本模型推理阶段输出的形状是[B, T, C]其中 T 是时间步数C 是类别数。每一帧都对应一个字符概率分布但相邻帧通常会预测同一个字符而且中间还会穿插 blank 帧。解码要做的事就是把这串概率序列换算成人类可读的字符串。最简单的解码方式是贪心每个时间步直接取概率最大的类别然后合并相邻重复字符、去掉 blank。贪心实现的 PyTorch 版本并不复杂关键是合并顺序。先合并连续重复再删除 blank两个顺序不能颠倒def greedy_decode(output, blank0): # output: [T, C] 概率张量log_softmax 之后 preds output.argmax(dim1).tolist() result [] prev None for p in preds: if p ! blank and p ! prev: result.append(p) prev p return result这段代码里prev负责记录上一个原始预测值。如果两个连续时间步都预测同一个字符p ! prev这个条件会把后者过滤掉实现去重。注意这里存在一个天然缺陷真实文本中如果出现连续重复字符比如 hello 里的 llCTC 路径上两个 l 之间必须插入一个 blank 才能被正确解码模型必须有足够强的输出倾向去生成这个 blank。这也是贪心解码在长文本上准确率会下降的根本原因。下面这张表对比了两种解码算法的定位解码方式原理准确率速度贪心搜索逐帧取最大概率中等极快Beam Search维护多条候选路径较高耗时随 beam 宽度增长前缀束搜索合并相同前缀再排序最高最慢4.2 用束搜索替代贪心代价与收益怎么平衡Beam Search 的核心是每步保留概率最高的 K 条路径而不是只留一条。但直接对 CTC 路径做 Beam Search 有个问题同一段文本会对应多条不同对齐路径如果不做前缀合并beam 里会塞满重复内容。因此实践中更常用的是前缀束搜索它把共享相同前缀的路径概率相加再去重排序。一个可运行的简化版本可以用 Python 的heapq实现但工程上我更建议直接调用torchaudio里的torchaudio.functional.rnnt_loss配套的解码器或者用pyctcdecode这个库它对语言模型融合的支持更好。如果只想在现有代码里快速提升准确率可以先试试增加带语言模型的二次打分而不是一上来就改解码算法。4.3 CharBound 数据如何辅助对齐验证源码里的testCharBound.mat不是训练必需但它对检测模型是否存在对而不准的问题很有用。每个样本的字符边界坐标可以画成一条水平轴把每个时间步预测的字符位置投影到这条轴上就能直观看到模型在哪几个时刻出现了跳变或重复。这个可视化用 matplotlib 即可完成import matplotlib.pyplot as plt # char_bounds: 每个字符的 [start, end] 坐标 # preds: 解码后的字符索引列表 fig, ax plt.subplots(figsize(10, 3)) for i, (s, e) in enumerate(char_bounds): ax.plot([s, e], [1, 1], linewidth4, labelfchar {i}) for t, p in enumerate(preds): ax.text(t * 10, 0, chars[p], fontsize8, hacenter) ax.set_yticks([]) plt.show()这段代码里char_bounds的坐标系必须和输入图片一致否则画出来的对应关系会偏移。一般我会先把真实边界画在图上再把模型预测的字符中心点画上去观察两者之间的偏移量。如果偏移保持一致说明模型学到了稳定的左到右阅读顺序如果偏移忽大忽小一般说明 CNN 部分提取的宽度特征不稳定需要检查输入图像是否做了端到端的归一化。5. 工程迁移把 CRNN 改造成验证码识别器与 TensorBoard 调优5.1 从 IIIT-5k 迁移到验证码数据改哪里验证码识别和场景文字识别的最大区别是字符集小、字符间距均匀、干扰线多。迁移时不需要改网络结构重点改三个地方字符表、图片预处理、输出类别数。假设验证码是 4 位数字那么n_classes就是 10 个数字加 1 个空白共 11 类而不是从原模型继承整个字典。预处理上要把验证码先转灰度再做二值化或去干扰。我给一个常用的预处理思路def preprocess_captcha(img_path, height32): img cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) _, img cv2.threshold(img, 127, 255, cv2.THRESH_BINARY_INV) h, w img.shape scale height / h img cv2.resize(img, None, fxscale, fyscale, interpolationcv2.INTER_CUBIC) img img.astype(np.float32) / 255.0 return imgTHRESH_BINARY_INV可以应对大多数白底深色文字但带噪声的验证码还需要配合形态学操作去孤立噪点。这个控制在数据量小的场景下能明显提升收敛速度。5.2 TensorBoard 里到底该看哪几条曲线源码的ctc_pytorch_tensorboard.ipynb文件名里直接带tensorboard说明作者在训练时就觉得文本损失不够直观。我的经验是除了记录train_loss一定要记录三个指标CTC loss 的滑动平均、字符准确率而非整串准确率、以及学习率实际生效值。写入的方式很简单from torch.utils.tensorboard import SummaryWriter writer SummaryWriter(log_dir./runs/crnn_captcha) # 每个 step 结束时 writer.add_scalar(loss/ctc, loss.item(), global_stepglobal_step) # 每个 epoch 结束时 writer.add_scalar(acc/char_acc, char_acc, global_stepepoch) writer.add_scalar(lr/current, optimizer.param_groups[0][lr], global_stepepoch) writer.close()lr/current这条曲线特别容易被忽略因为有些代码虽然scheduler.step()写了但 StepLR 的gamma设置后不会在训练日志里自动体现。我当时排查过一次loss 下降变慢的问题最后发现是学习率在 epoch 10 后跌到了1e-5但代码里没人察觉。5.3 训练不收敛时先查这三个地方如果迁移后 loss 完全不下降先检查 blank 索引是否和字符表错位。统一约定字符映射从 1 开始0 永远留给 blank然后把CTCLoss(blank0)写死。第二件事是检查输入图像的高度是否真的是 32很多人把验证码图片直接 resize 成(32, 128)但没确认原始图片本身不是(28, 128)这样 CNN 的池化层会把高度压成负数维度的边界值训练直接崩到 nan。最后再查target_lengths是否小于input_lengths // 4CTC 在序列长度小于标签长度时会出现空洞梯度表现是 loss 偶尔跳成 inf加zero_infinityTrue只能治标真正治本还是要调大fix_width或减小 batch 内的最大文本长度。本文还有配套的精品资源点击获取