简介:一份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_size=256): super().__init__() # CNN 部分:把单通道灰度图映射成高层视觉特征 self.cnn = nn.Sequential( nn.Conv2d(1, 64, 3, padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(2, 2), # 高度、宽度各减半 nn.Conv2d(64, 128, 3, padding=1), nn.ReLU(inplace=True), nn.MaxPool2d(2, 2), # 高度、宽度再减半 nn.Conv2d(128, 256, 3, padding=1), nn.BatchNorm2d(256), nn.ReLU(inplace=True), nn.Conv2d(256, 256, 3, padding=1), nn.ReLU(inplace=True), nn.MaxPool2d((2, 1)), # 只压缩高度,保留宽度 nn.Conv2d(256, 512, 3, padding=1), nn.BatchNorm2d(512), nn.ReLU(inplace=True), nn.Conv2d(512, 512, 3, padding=1), nn.ReLU(inplace=True), nn.MaxPool2d((2, 1)), # 只压缩高度 nn.Conv2d(512, 512, 3, padding=1), nn.BatchNorm2d(512), nn.ReLU(inplace=True), ) # RNN 部分:双向 LSTM 建模序列上下文 self.lstm = nn.LSTM(512, hidden_size, num_layers=2, bidirectional=True, batch_first=True) 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(blank=0, zero_infinity=True) # 假设 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 # 每个样本的标签长度 )blank=0表示类别索引 0 被保留给空字符,这也意味着字符表构建时要从 1 开始编号。zero_infinity=True是个容易被忽略的细节:当某个 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.pkl
3.1 源码里这些文件分别承担什么角色
第一次打开这个项目时,先别急着跑训练,把数据文件之间的关系理清楚能省很多调试时间。traindata.mat和testdata.mat存的是 IIIIT-5k 数据集的合成图片矩阵,trainCharBound.mat和testCharBound.mat存的是每个字符在图片中的边界框列表。训练时真正喂给模型的其实是train_fix_width.pkl,它把图片统一处理成了固定宽度,并和标签索引一一对应。
文件和用途对照:
| 文件 | 内容 | 用途 |
|---|---|---|
traindata.mat | 训练图像矩阵 | 原始数据,需要转成 png 或 npy |
testdata.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_height=32, fix_width=280): 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)), mode='bilinear' ).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], dim=2) else: img = img[:, :, :self.fix_width] target = torch.tensor([self.char_dict[c] for c in label], dtype=torch.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], dtype=torch.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_classes=len(char_dict) + 1) optimizer = torch.optim.Adam(model.parameters(), lr=1e-3) scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=10, gamma=0.1) loader = DataLoader(dataset, batch_size=32, shuffle=True, collate_fn=collate_ocr, drop_last=True) 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, dtype=torch.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_size=10配合gamma=0.1是 OCR 任务里常见的阶梯式下降策略,epoch 10 之前模型还在学习字符的局部纹理,过早降学习率会让 RNN 部分难以学到长距离依赖。drop_last=True是为了防止最后一个 batch 样本数过少导致input_lengths和target_lengths出现极端值。训练过程中如果发现 loss 一直在 1 附近震荡,优先检查字符表里是否混入了重复字符或全角符号,这类噪声会让 CTC 的概率分布长期无法收敛到单峰。
4. 推理解码:贪心搜索、束搜索与字符边界还原
4.1 解码的本质是从概率分布中恢复文本
模型推理阶段输出的形状是[B, T, C],其中 T 是时间步数,C 是类别数。每一帧都对应一个字符概率分布,但相邻帧通常会预测同一个字符,而且中间还会穿插 blank 帧。解码要做的事就是把这串概率序列换算成人类可读的字符串。最简单的解码方式是贪心:每个时间步直接取概率最大的类别,然后合并相邻重复字符、去掉 blank。
贪心实现的 PyTorch 版本并不复杂,关键是合并顺序。先合并连续重复,再删除 blank,两个顺序不能颠倒:
def greedy_decode(output, blank=0): # output: [T, C] 概率张量(log_softmax 之后) preds = output.argmax(dim=1).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" 里的 "ll",CTC 路径上两个 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], linewidth=4, label=f'char {i}') for t, p in enumerate(preds): ax.text(t * 10, 0, chars[p], fontsize=8, ha='center') 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, height=32): 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, fx=scale, fy=scale, interpolation=cv2.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_step=global_step) # 每个 epoch 结束时 writer.add_scalar('acc/char_acc', char_acc, global_step=epoch) writer.add_scalar('lr/current', optimizer.param_groups[0]['lr'], global_step=epoch) writer.close()lr/current这条曲线特别容易被忽略,因为有些代码虽然scheduler.step()写了,但 StepLR 的gamma设置后不会在训练日志里自动体现。我当时排查过一次"loss 下降变慢"的问题,最后发现是学习率在 epoch 10 后跌到了1e-5,但代码里没人察觉。
5.3 训练不收敛时先查这三个地方
如果迁移后 loss 完全不下降,先检查 blank 索引是否和字符表错位。统一约定:字符映射从 1 开始,0 永远留给 blank,然后把CTCLoss(blank=0)写死。第二件事是检查输入图像的高度是否真的是 32,很多人把验证码图片直接 resize 成(32, 128),但没确认原始图片本身不是(28, 128),这样 CNN 的池化层会把高度压成负数维度的边界值,训练直接崩到 nan。最后再查target_lengths是否小于input_lengths // 4,CTC 在序列长度小于标签长度时会出现空洞梯度,表现是 loss 偶尔跳成 inf,加zero_infinity=True只能治标,真正治本还是要调大fix_width或减小 batch 内的最大文本长度。
本文还有配套的精品资源,点击获取