目录

目录

前言

CRNN介绍

一、网络结构

1.1 概述

CNN(卷积神经网络)

调整输入形状

LSTM(长短时记忆网络)

1.2 代码实现

CNN部分

LSTM部分

CRNN主模块

1.3 模型完整代码

二.训练

2.1 数据处理

词表

CrnnDataSet

load_data

1. 加载原始数据并随机打乱

2.划分训练/验证索引

3. 初始化字符映射表

4. 构建完整数据集

5.利用Subset划分训练集和验证集

6.创建 DataLoader

数据集推荐

2.2 训练

main()

1. 设备设置与数据加载

2. 模型初始化

3. 检查点路径与训练起始设置

4. 训练组件配置

5. 从检查点恢复训练

6. 训练主循环

7. 每轮训练与验证

8. 模型保存与可视化

train()

1. 启用训练模式

2. 初始化损失累加器

3. 遍历训练数据

4. 数据迁移到设备 & 梯度清零

5. 前向传播

6. 构造 CTC 损失所需的输入长度

7. 安全检查:防止目标长度超过输出长度

8. 计算 CTC 损失

9. 混合精度反向传播与优化

10. 损失记录与进度打印

11. 返回平均损失

validate()

2.3 完整代码

2.3.1 数据处理类

2.3.2 工具文件

2.3.3 训练文件

三、推理

四、总结


前言

CRNN介绍

关于CRNN,博主也在不断学习中,也是有点懵懵懂懂,但是其大致可以分为三个部分,CNN+LSTM+CTC损失函数,本文主要分享怎么代码实现。

在实际项目中,我们经常会遇到需要识别图像中文字序列的任务,比如车牌识别、身份证信息提取、自然场景文字识别(Scene Text Recognition)等。这类任务的难点在于:文字长度不固定、字体多样、背景复杂,传统的分类模型难以直接处理。

为了解决这一问题,CRNN(Convolutional Recurrent Neural Network) 应运而生。它结合了卷积神经网络(CNN)、循环神经网络(RNN/LSTM)和连接时序分类(CTC)损失函数,能够端到端地处理不定长的文本图像识别任务。

随着技术的发展,新兴的基于transformers(VIT)的视觉模型也开始崭露头角。

上图是博客『带你学AI』一文带你搞懂OCR识别算法CRNN:解析+源码-CSDN博客中CRNN的结构图,可以很直观的看到CRNN分成了三大模块,CNN卷积神经网络+RNN循环神经网络以及最后的CTC loss。

一、网络结构

1.1 概述

CNN(卷积神经网络)

在CNN部分,我们通过多个卷积层和下采样层逐步提取特征并降低空间维度。例如,输入图像形状为(1, 1, 32, 128)(batch×通道数×高度×宽度)。如果CNN部分包含5个卷积层,通常会这样设计:

卷积后向下池化,直到将高度维度从32降到1

第一层:(1, 1,32,128)-> (1, C1, 16,64)

第二层:(1, C1, 16,64)-> (1, C2, 8,32)

第三层:(1, C2,  8,32) - > (1, C3, 4,32)

第四层:(1, C3, 4,32) - > (1, C4,  2,32)

第五层:(1, C4, 2,32) - > (1, C5,  1,32)

注意:从第3层开始,池化只在高度方向进行(池化核为2×1),这样宽度(32)保持不变,可以相对保留图片在宽度维度上的信息。

至于为什么要这么做,博主讲一下自己的理解:

在处理像文本图像这样具有明确方向性的数据时(比如从左到右书写的文字),我们通常会把卷积神经网络(CNN)提取出的特征图,在宽度方向上展开成一个序列。这是因为文字的阅读顺序是从左到右的,宽度方向天然就包含了这种时序或顺序信息

而高度方向则不同——在CNN的多层卷积过程中,上下方向的空间信息(比如字符的形状、笔画结构等)已经被充分提取并压缩了,所以不需要再保留高度维度的序列信息。

接下来,为了把特征输入到擅长处理时序数据的LSTM中,我们需要让数据符合LSTM的输入格式,通常是(序列长度, 批次大小, 特征维度)。因此,通常会把特征图的高度维度压缩成1(比如通过池化或直接reshape),只保留宽度方向作为序列长度,这样就能顺利送入LSTM进行后续的序列建模了

调整输入形状

上文说到LSTM的输入是(序列长度, 批次大小, 特征维度),但从cnn中输出的形状是(1, C5,  1,32),因此我们需要丢弃高度维度(上文中也说到在CNN的多层卷积过程中,上下方向的空间信息已经被充分提取,并且高度已经被压缩为1,所以丢弃没有影响),此时的形状为(1,c5,32)

接着为了与 LSTM 的输入格式对齐,接下来我们需要将这个张量重新组织:把宽度维度(32)视为序列长度,通道维度(C5)作为每个时间步的特征维度,而第一个维度(1)对应批次大小。于是,通过适当的转置和 reshape 操作,就能将其转换为 LSTM 所需的(32, 1, C5)格式,顺利送入后续的时序建模模块。

LSTM(长短时记忆网络)

因为pytroch有一个封装好了的LSTM模块,所以这里没有重新实现(ps:就是懒😅)
这里使用的是一个双向LSTM,共堆叠了四层。
为什么使用双向的,这是因为双向可以同时从前向和后向两个方向处理数据。

这里举一个例子,识别一个单词如'hello':

  • 当只使用前向LSTM 时,假设输入图像质量较差,其中字母 "e" 因模糊或光照问题,看起来很像 "c""o"。此时,仅靠局部特征很难判断中间这个字符到底是哪个,它可能会输出hcllo或hollo等。

  • 而使用双向LSTM 时,模型在处理该位置时,不仅能“看到”前面的 "h",还能“看到”后面的 "l", "l", "o"(感觉和BERT很像)。于是模型会推理:“如果前面是 'h',后面是 'llo',那么中间这个模糊字符极大概率是 'e',因为 'hello' 是一个高频单词,而 'hcllo' 或 'hollo' 几乎不存在(或极少见)。”
    (ps:这里的推理指的是根据前后字母计算的为e的概率)

总结一下,为什么使用双向lstm:

在 OCR或文本图像识别任务中,模型的目标是对整张图像中的文本序列进行一次性判别(也就是输入是完整的图像信息,输出是对应每个位置的字符预测),而不是像语言模型那样逐字生成。模型可以利用整个序列的信息来做每个位置的预测,而双向结构天然契合这种“全局上下文感知”的需求。

到现在,基本模型网络结构已经搭建完成了,为了防止训练的时候梯度消失和提升模型的性能,在cnn和LSTM部分都加入了残差连接。

1.2 代码实现

下面则是具体实现

CNN部分
class CNNFeatureExtractor(nn.Module):
    def __init__(self, in_channels):
        super(CNNFeatureExtractor, self).__init__()
        ##第一层卷积核
        self.conv1 = nn.Sequential(
            nn.Conv2d(in_channels, 32, kernel_size=3, stride=1, padding=1),
            nn.BatchNorm2d(32),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2, stride=2) #16,64
        )
        self.conv2 = nn.Sequential(
            BasicConvBlock(32,64,stride=2, downsample=nn.Sequential(
                nn.Conv2d(32, 64, kernel_size=1, stride=2, padding=0),
                nn.BatchNorm2d(64)
            )),
            BasicConvBlock(64,64,stride=1)
        ) ##8,32
        self.conv3 = nn.Sequential(
            BasicConvBlock(64,128,stride=(2,1), downsample=nn.Sequential(
                nn.Conv2d(64, 128, kernel_size=1, stride=(2,1), padding=0),
                nn.BatchNorm2d(128)
            )),
            BasicConvBlock(128,128,stride=1)
        ) ##4,32
        self.conv4 = nn.Sequential(
           BasicConvBlock(128,256,stride=(2,1), downsample=nn.Sequential(
                nn.Conv2d(128, 256, kernel_size=1, stride=(2,1), padding=0),
                nn.BatchNorm2d(256)
            )),
            BasicConvBlock(256,256,stride=1)
        ) ##2,32
        self.conv5 = nn.Sequential(
            BasicConvBlock(256,256,stride=1),
            BasicConvBlock(256,256,stride=1)
        )
        self.conv6 = nn.Sequential(
            BasicConvBlock(256,512,stride=(2,1), downsample=nn.Sequential(
                nn.Conv2d(256, 512, kernel_size=1, stride=(2,1), padding=0),
                nn.BatchNorm2d(512)
            )),
            BasicConvBlock(512,512,stride=1)
        ) ##1,32

    def forward(self, x):
        x = self.conv1(x)
        x = self.conv2(x)
        x = self.conv3(x)
        x = self.conv4(x)
        x = self.conv5(x)
        x = self.conv6(x)
        x = nn.AdaptiveAvgPool2d((1, x.shape[-1]))(x)
        x = x.squeeze(2)   ##(b, c, h, w)
        x = x.permute(2,0,1)
        return x

其中BasicConvBlock是定义的残差块,其网络结构如下

class BasicConvBlock(nn.Module):
    def __init__(self, in_channels, out_channels, stride, downsample=None):
        super(BasicConvBlock, self).__init__()
        self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1, bias=False)
        self.bn = nn.BatchNorm2d(out_channels)
        self.relu = nn.ReLU()

        self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False)
        self.bn2 = nn.BatchNorm2d(out_channels)

        self.downsample = downsample
        self.stride = stride

    def forward(self, x):
        residual = x
        out = self.conv(x)
        out = self.bn(out)
        out = self.relu(out)
        out = self.conv2(out)
        out = self.bn2(out)
        if self.downsample:
            residual = self.downsample(x)
        out += residual
        out = self.relu(out)
        return out

下面将详细解释每段代码的意义和作用

CNNFeatureExtractor(cnn模块)
要看这段代码是什么作用,我一般是先看forward方法里面做了什么计算,那如上述的CNNFeatureExtractor的forward里面,输入张量x会依次通过 conv1到conv6六个卷积阶段,每一阶段都对特征图的空间维度(高和宽)进行不同程度的下采样,同时逐步增加通道数,从而提取更高层次、更抽象的特征表示(这里对应概述中对图像进行不断卷积和下采样的部分)。

当从最后一个卷积块中输出的时候,x的形状就是(batch, 通道数,1,32)

继续,

x = nn.AdaptiveAvgPool2d((1, x.shape[-1]))(x)

这段代码是为了确保特征图的高度(H)被强制压缩为 1,而保留原始的宽度不变,从而得到一个形状为 (B, C, 1, W) 的张量。详细解释一下,AdaptiveAvgPool2d方法的作用是将任意尺寸的输入特征图,转换为指定的输出尺寸,(1x.shape[-1])表示高度压缩为1,宽度还是这里的32,(x)表示调用这个平均池化的对象,等价于:

pool = nn.AdaptiveAvgPool2d((1, x.size(-1)))
x = pool(x)

再继续,后面的

x = x.squeeze(2)   
x = x.permute(2,0,1)

则是概述中的调整输入形状部分,首先通过 x.squeeze(2) 移除第2维(即高度维度)中大小为1的维度,然后使用 x.permute(2, 0, 1) 对维度进行重新排列。

,在forward方法走完后,回到最开始的

x = self.conv1(x)
...

部分,这里我们将逐层分析每一层卷积块的作用

记住x最开始的形状是(1,1,32,128)

第一层卷积

self.conv1 = nn.Sequential(
    nn.Conv2d(in_channels, 32, kernel_size=3, stride=1, padding=1),
    nn.BatchNorm2d(32),                                              
    nn.ReLU(),                                                       
    nn.MaxPool2d(kernel_size=2, stride=2)   
)                               
  • 首先通过一个 3×3 卷积将输入通道映射到 32 维,in_channels输出通道数(也就是1),32是输出通道数,kernel_size卷积核大小为3,stride步长为1,即卷积核每次只移动1个像素,padding填充为1,至于为什么要填充为1,因为如下公式,input_size和output_size都是H,可以算出padding为1      

  • 归一化

  • 激活函数ReLU

  • 接着使用 MaxPool2d 池化函数对高和宽同时进行 2 倍下采样(即 H → H/2W → W/2,这里就是32→16,128→64

现在的形状是(1,32,16,64)

第二层到第六层

self.conv2 = nn.Sequential(
    BasicConvBlock(32,64,stride=2, downsample=nn.Sequential(
        nn.Conv2d(32, 64, kernel_size=1, stride=2, padding=0),
        nn.BatchNorm2d(64)
    )),
    BasicConvBlock(64,64,stride=1)
)

conv2 开始,网络采用 残差块(BasicConvBlock) 构建,这是借鉴 ResNet 的设计思想,有助于缓解深层网络中的梯度消失问题,并提升特征表达能力。con2~con6网络结构都是一样的,只是传入参数不一样。

这里暂时不介绍BasicConvBlock残差块,防止逻辑混乱。至于为什么要在每个Sequential里面堆叠两个BasicConvBlock,是参照的ResNet-18的网络结构,目的是第一个 block:负责“维度变换”,第二个 block:负责“特征精炼”

下面列出的是经过每一层残差块后的形状

conv2 -> (1, 64, 8, 32)

conv3 -> (1, 128, 4, 32)

conv4 -> (1, 256, 2, 32)

conv5 -> (1, 256, 2, 32)

conv6 -> (1, 512, 1, 32)

刚刚好对应(batch, 通道数,1,32)


BasicConvBlock介绍

同样还是看forward,计算过程如下:

第一层卷积->归一化->激活函数->第二层卷积->归一化

第一层卷积下采样,第二层卷积用来提取特征

说明一下,为什么这里的没有使用池化,这是因为下采样已通过卷积层的 stride=2 实现,无需额外使用池化层。卷积的stride=2本身就可以完成下采样,效果和使用MaxPool2d(stride=2)一致,而卷积的下采样因为参数可以学习,于是下采样的同时还能学习特征,提升模型的表达能力,下图是输入5*5,卷积核3*3,stride = 2,padding = 0的动态效果

if self.downsample:
    residual = self.downsample(x)
我们先在这里写出downsample是个啥东西:
nn.Sequential(
    nn.Conv2d(32, 64, kernel_size=1, stride=2, padding=0),
    nn.BatchNorm2d(64)
)

可以看到downsample内部是一个stride为2的下采样卷积层,为什么要使用downsample,我们可以看最开始定义的

residual = x
以及下方的
out += residual

这是残差连接的具体表现,out在经过一系列卷积下采样后形状变成了(batch,通道数,w/2,h/2)

而未经过downsample的residual的形状却是(batch,通道数,w,h),此时两者相加会因为两个张量的 shape 不匹配而报错。经过downsample后,residual也向下采样,形状变为(batch,通道数,w/2,h/2),这样就可以和out相加了

将残差后得到的值在进行一次激活就可以作为整个残差块的输出了


LSTM部分
class LSTMFeatureExtractor(nn.Module):
    def __init__(self, input_size, hidden_size, num_layers, bidirectional, output_size):
        super(LSTMFeatureExtractor, self).__init__()
        self.layers = nn.ModuleList()
        current_size = input_size
        for i in range(num_layers):
            self.layers.append(LSTMConvBlock(current_size, hidden_size, bidirectional))
            current_size = hidden_size * (2 if bidirectional else 1)
        self.dropout = nn.Dropout(0.3)
        self.fnn1 = nn.Linear(current_size, 256)
        self.fnn2 = nn.Linear(256, output_size)


    def forward(self, x):
        for layer in self.layers:
            x = layer(x)
        x = self.dropout(x)
        x = self.fnn1(x)
        x = torch.relu(x)
        x = self.fnn2(x)
        x = torch.relu(x)
        return x
LSTMConvBlock
class LSTMConvBlock(nn.Module):
    def __init__(self, input_size, hidden_size, bidirectional):
        super(LSTMConvBlock, self).__init__()
        self.bidirectional = bidirectional
        self.layer_norm = nn.LayerNorm(input_size)
        self.lstm = nn.LSTM(input_size, hidden_size, num_layers=1, bidirectional=bidirectional)
        output_size = hidden_size * (2 if bidirectional else 1)
        if output_size != input_size:
            self.fc = nn.Linear(input_size, output_size)
        else:
            self.fc = None

    def forward(self, x):
        residual = x
        x = self.layer_norm(x)
        x, _ = self.lstm(x)
        if self.fc is not None:
            residual = self.fc(residual)
        x = x + residual
        return x

在LSTM部分,为了防止因为lstm层数过多而引发的梯度消失问题,也使用了残差连接

还是先看主模块LSTMFeatureExtractor的forward()方法

for layer in self.layers:
    x = layer(x)

这里的self.layers是一个堆叠了num_layers=4层的数组,也就是说lstm部分共有四层,注意看self.layers的实现:

for i in range(num_layers):
    self.layers.append(LSTMConvBlock(current_size, hidden_size, bidirectional))
    current_size = hidden_size * (2 if bidirectional else 1)

其中:

  • num_layers=4,表示总共堆叠了 4 层 LSTM;
  • 每一层都是一个 LSTMConvBlock,封装了 LSTM 层(可选双向)。
  • current_size 动态更新为下一层的输入维度:若使用双向 LSTM,则输出维度为 2 * hidden_size,否则为 hidden_size

后续经过dropout丢弃一部分参数,在经过两层全连接层提升模型泛化能力

LSTMConvBlock

self.lstm = nn.LSTM(input_size, hidden_size, num_layers=1, bidirectional=bidirectional)

这句话的意思是创建一个输入为input_size,输出为hidden_size的单层的双向LSTM层

来看forward方法

先是对x进行了归一化操作,稳定训练过程,加速收敛

再传入LSTM层

而后面的

if self.fc is not None:
    residual = self.fc(residual)

作用也是为了对齐out和residual的维度,使计算不会报错

CRNN主模块
class CRNNModel(nn.Module):
    def __init__(self, in_channels=1, num_classes=35, hidden_size=512, num_layers=4, bidirectional=True, output_size=128):
        super(CRNNModel, self).__init__()
        self.cnn = CNNFeatureExtractor(in_channels)
        self.lstm = LSTMFeatureExtractor(512, hidden_size, num_layers, bidirectional, output_size)
        self.dropout = nn.Dropout(0.5)
        self.fnn = nn.Sequential(
            nn.Linear(output_size, 256),
            nn.ReLU(),
            nn.Dropout(0.3),
            nn.Linear(256, 512),
            nn.ReLU(),
            nn.Linear(512, num_classes)
        )

    def forward(self, x):
        x = self.cnn(x)
        x = self.lstm(x)
        x = self.dropout(x)
        x = self.fnn(x)
        return x

该模块就是整合了CNN和LSTM,同时添加了一个分类头fnn

1.3 模型完整代码

文件Model:

import torch
from torch import nn


class BasicConvBlock(nn.Module):
    def __init__(self, in_channels, out_channels, stride, downsample=None):
        super(BasicConvBlock, self).__init__()
        self.conv = nn.Conv2d(in_channels, out_channels, kernel_size=3, stride=stride, padding=1, bias=False)
        self.bn = nn.BatchNorm2d(out_channels)
        self.relu = nn.ReLU()

        self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=3, stride=1, padding=1, bias=False)
        self.bn2 = nn.BatchNorm2d(out_channels)

        self.downsample = downsample
        self.stride = stride

    def forward(self, x):
        residual = x
        out = self.conv(x)
        out = self.bn(out)
        out = self.relu(out)
        out = self.conv2(out)
        out = self.bn2(out)
        if self.downsample:
            residual = self.downsample(x)
        out += residual
        out = self.relu(out)
        return out

class CNNFeatureExtractor(nn.Module):
    def __init__(self, in_channels):
        super(CNNFeatureExtractor, self).__init__()
        ##第一层卷积核
        self.conv1 = nn.Sequential(
            nn.Conv2d(in_channels, 32, kernel_size=3, stride=1, padding=1),
            nn.BatchNorm2d(32),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2, stride=2) #16,64
        )
        self.conv2 = nn.Sequential(
            BasicConvBlock(32,64,stride=2, downsample=nn.Sequential(
                nn.Conv2d(32, 64, kernel_size=1, stride=2, padding=0),
                nn.BatchNorm2d(64)
            )),
            BasicConvBlock(64,64,stride=1)
        ) ##8,32
        self.conv3 = nn.Sequential(
            BasicConvBlock(64,128,stride=(2,1), downsample=nn.Sequential(
                nn.Conv2d(64, 128, kernel_size=1, stride=(2,1), padding=0),
                nn.BatchNorm2d(128)
            )),
            BasicConvBlock(128,128,stride=1)
        ) ##4,32
        self.conv4 = nn.Sequential(
           BasicConvBlock(128,256,stride=(2,1), downsample=nn.Sequential(
                nn.Conv2d(128, 256, kernel_size=1, stride=(2,1), padding=0),
                nn.BatchNorm2d(256)
            )),
            BasicConvBlock(256,256,stride=1)
        ) ##2,32
        self.conv5 = nn.Sequential(
            BasicConvBlock(256,256,stride=1),
            BasicConvBlock(256,256,stride=1)
        )
        self.conv6 = nn.Sequential(
            BasicConvBlock(256,512,stride=(2,1), downsample=nn.Sequential(
                nn.Conv2d(256, 512, kernel_size=1, stride=(2,1), padding=0),
                nn.BatchNorm2d(512)
            )),
            BasicConvBlock(512,512,stride=1)
        ) ##1,32

    def forward(self, x):
        x = self.conv1(x)
        x = self.conv2(x)
        x = self.conv3(x)
        x = self.conv4(x)
        x = self.conv5(x)
        x = self.conv6(x)
        x = nn.AdaptiveAvgPool2d((1, x.shape[-1]))(x)
        x = x.squeeze(2)
        x = x.permute(2,0,1)
        return x


class LSTMConvBlock(nn.Module):
    def __init__(self, input_size, hidden_size, bidirectional):
        super(LSTMConvBlock, self).__init__()
        self.bidirectional = bidirectional
        self.layer_norm = nn.LayerNorm(input_size)
        self.lstm = nn.LSTM(input_size, hidden_size, num_layers=1, bidirectional=bidirectional)
        output_size = hidden_size * (2 if bidirectional else 1)
        if output_size != input_size:
            self.fc = nn.Linear(input_size, output_size)
        else:
            self.fc = None

    def forward(self, x):
        residual = x
        x = self.layer_norm(x)
        x, _ = self.lstm(x)
        if self.fc is not None:
            residual = self.fc(residual)
        x = x + residual
        return x

class LSTMFeatureExtractor(nn.Module):
    def __init__(self, input_size, hidden_size, num_layers, bidirectional, output_size):
        super(LSTMFeatureExtractor, self).__init__()
        self.layers = nn.ModuleList()
        current_size = input_size
        for i in range(num_layers):
            self.layers.append(LSTMConvBlock(current_size, hidden_size, bidirectional))
            current_size = hidden_size * (2 if bidirectional else 1)
        self.dropout = nn.Dropout(0.3)
        self.fnn1 = nn.Linear(current_size, 256)
        self.fnn2 = nn.Linear(256, output_size)


    def forward(self, x):
        for layer in self.layers:
            x = layer(x)
        x = self.dropout(x)
        x = self.fnn1(x)
        x = torch.relu(x)
        x = self.fnn2(x)
        x = torch.relu(x)
        return x

class CRNNModel(nn.Module):
    def __init__(self, in_channels=1, num_classes=35, hidden_size=512, num_layers=4, bidirectional=True, output_size=128):
        super(CRNNModel, self).__init__()
        self.cnn = CNNFeatureExtractor(in_channels)
        self.lstm = LSTMFeatureExtractor(512, hidden_size, num_layers, bidirectional, output_size)
        self.dropout = nn.Dropout(0.5)
        self.fnn = nn.Sequential(
            nn.Linear(output_size, 256),
            nn.ReLU(),
            nn.Dropout(0.3),
            nn.Linear(256, 512),
            nn.ReLU(),
            nn.Linear(512, num_classes)
        )

    def forward(self, x):
        x = self.cnn(x)
        x = self.lstm(x)
        x = self.dropout(x)
        x = self.fnn(x)
        return x









二.训练

2.1 数据处理

词表

首先得加载词表,这里使用的是paddleOCR的词表ppocr_keys_v1.txt

def init_chat_to_idx():
    chat_to_idx = {
        '<blank>': 0,
    }
    idx = len(chat_to_idx)
    with open('ppocr_keys_v1.txt', 'r', encoding='utf-8-sig') as f:
        for line in f:
            chat = line.strip()
            chat_to_idx[chat] = idx
            idx += 1
    idx_to_char = {v: k for k, v in chat_to_idx.items()}
    return chat_to_idx, idx_to_char

chat_to_idx = { '<blank>': 0}解释:在连接时序分类(CTC)中,<blank> 是一个关键概念,用于解决输入输出不对齐的问题,也就是说后续的计算CTC loss时<blank>很重要

CrnnDataSet
CrnnDataSet是继承了Dataset的一个自定义数据集处理的类

其代码如下

class CrnnDataSet(Dataset):
    def __init__(self, data_path, labels, chat_to_idx, transform=None):
        self.data_path = data_path
        self.labels = labels
        self.chat_to_idx = chat_to_idx
        self.transform = transform or transforms.Compose([
            transforms.Grayscale(),
            transforms.ToTensor(),
            transforms.Normalize((0.1,), (0.3,))
        ])

    def __len__(self):
        return len(self.data_path)

    def __getitem__(self, idx):
        img_path = self.data_path[idx]
        label = self.labels[idx]
        img = self.preprocess_image(img_path)
        filtered_label = []
        for c in label:
            if c in self.chat_to_idx:
                filtered_label.append(self.chat_to_idx[c])
        return {'image': img, 'label': torch.tensor(filtered_label, dtype=torch.long)}

    def preprocess_image(self,img_path):
        img = Image.open(img_path)
        img = self.transform(img)
        return img

这里从__getitem__ 方法开始,__getitem__ 是 PyTorch Dataset 的核心方法。当使用 DataLoader 加载数据时,它会调用此方法,根据索引 idx 返回一个样本(图像 + 标签)。

步骤如下

获取图像路径和标签

img_path = self.data_path[idx]
label = self.labels[idx]
  • data_path 是一个列表,包含所有图像的路径(如 ['/data/001.png', '/data/002.png', ...])。
  • labels也 是一个列表,每个元素是对应图像的文本标签(如 ['hello', '你好', ...])。

图像预处理

调用 preprocess_image 方法,将图像路径转换为标准化的张量(Tensor)。

img = self.transform(img)中的transform部分作用:将图片转换为灰度图,转为Tensor张量

标签字符映射为索引

filtered_label = []
for c in label:
    if c in self.chat_to_idx:
        filtered_label.append(self.chat_to_idx[c])
  • chat_to_idx 是由上文中的init_chat_to_idx返回的一个字典,例如:{'a': 1, 'b': 2, ..., '<blank>':: 0}
  • 将文本标签中的每个字符转换为对应的数字索引。
  • 注意:如果某个字符不在 chat_to_idx 中(比如生僻字或标点),会被直接跳过(过滤掉)。这可以防止模型崩溃,但也可能导致信息丢失。实际项目中,建议加入 <UNK>(未知字符)处理。
load_data
def load_data(file_path, file_root):
    imgs, labels = utils.load_data(file_path, file_root)
    indices = list(range(len(imgs)))
    train_indices, val_indices = train_test_split(
        indices, test_size=0.01, random_state=42
    )
    chat_to_idx, idx_to_chat = utils.init_chat_to_idx()
    full_dataset = CrnnDataSet(imgs, labels, chat_to_idx)
    train_dataset = Subset(full_dataset, train_indices)
    val_dataset = Subset(full_dataset, val_indices)
    print(f"训练集大小: {len(train_dataset)}, 验证集大小: {len(val_dataset)}")
    train_dataloader = DataLoader(train_dataset, batch_size=128, shuffle=True, collate_fn=collate_fn, num_workers=4)
    val_dataloader = DataLoader(val_dataset, batch_size=128, shuffle=True, collate_fn=collate_fn, num_workers=4)
    return train_dataloader, val_dataloader, chat_to_idx, idx_to_chat

在训练时,我们需要将原始数据划分为训练集和验证集,并封装成 PyTorch 的 DataLoader,以便高效地送入模型训练。

步骤解析:

1. 加载原始数据并随机打乱

这里的utils.load()就是使用

with open(file_path, 'r', encoding='c') as f

去数据集文件夹中加载图片路径和对应的标签

2.划分训练/验证索引
  • 使用 sklearn.model_selection.train_test_split 随机划分索引。
  • test_size=0.01 表示保留 1% 作为验证集(适用于大数据集)。
  • random_state=42 确保每次运行划分结果一致,便于复现。
3. 初始化字符映射表

调用初始化词表方法init_chat_to_idx返回词表字典

4. 构建完整数据集

使用前面定义的 CrnnDataSet 类,将原始数据封装为 PyTorch 数据集对象。

5.利用Subset划分训练集和验证集
6.创建 DataLoader

这里需要座着重介绍一下collate_fn函数,这个函数是自定义的用来来把一批(batch)样本“堆叠”(stack)成张量(Tensor)的函数。为什么要自定义呢,因为PyTorch 的 DataLoader 默认会尝试将一个 batch 中的所有样本“堆叠”成张量。而在该任务中,每张图片的宽度和高度都不一样,其对应的标签长度也不一样,如果使用默认的collate_fn函数会报错,于是自定义的collate_fn函数要实现以下目标:统一图像尺寸、保留标签变长特性,并为 CTC Loss 准备所需格式

目标有了那就可以实现了,但是在此之前,还要详细介绍统一图像尺寸、保留标签变长特性,并为 CTC Loss 准备所需格式。是什么意思.

统一图像尺寸: 需统一为相同尺寸(例如上文中的模型输入的高度就是固定为 32,宽度可变但 在同一batch 内需要对齐,一般是将该batch内的所有图片宽度填充到该batch内最宽图片的宽度)

CTC Loss 输入格式:

log_probs:模型输出

targets:拼接后的标签

input_length: 模型输出序列的时间步长度

targets_length:每个样本的标签序列的长度

具体实现如下:

def collate_fn(batches):
    images = []
    labels = []
    for batch in batches:
        images.append(batch['image'])
        labels.append(batch['label'])
    target_h = 32
    new_widths = []
    resized_images = []
    for img in images:
        c, h, w = img.shape
        scale = h / target_h
        new_width = int(w * scale)
        img_resized = F.resize(img, [target_h, new_width],antialias=True)
        resized_images.append(img_resized)
        new_widths.append(new_width)
    max_width = max(new_widths)
    padding_images = []
    for img in resized_images:
        pad_w = max_width - img.shape[-1]
        padding_img = functional.pad(img, (0, pad_w, 0, 0), 'constant', 0)
        padding_images.append(padding_img)
    images = torch.stack(padding_images, dim=0)
    targets = torch.cat(labels, dim=0)
    target_lengths = torch.LongTensor([len(label) for label in labels])
    return images, targets, target_lengths

接下来详细解释为什么要这样实现以及每行代码对应什么意思

  1. 图像与标签的初步收集
    首先,函数遍历输入的 batches(通常是由 DataLoader 传入的一组样本),分别提取每个样本中的图像(batch['image'])和标签(batch['label']),并分别存入 imageslabels 列表中。

  2. 等比例缩放图像高度至固定值(32)
    为了在保持图像宽高比的前提下统一高度,设定目标高度 target_h = 32。对每张图像,根据其原始高度 h 计算缩放比例 scale = h / target_h,并据此计算缩放后的宽度 new_width = int(w * scale)。随后使用 torchvision.transforms.functional.resize(即 F.resize)将图像等比缩放到 [target_h, new_width] 的尺寸,并启用抗锯齿(antialias=True)以提升缩放质量。缩放后的图像被存入 resized_images,其新宽度记录在 new_widths 中。

  3. 统一图像宽度:填充至批次内最大宽度
    由于不同图像缩放后的宽度可能不同,无法直接堆叠为一个四维张量(batch × channel × height × width)。为此,函数找出当前批次中所有图像的最大宽度 max_width = max(new_widths),然后对每张图像在其右侧(宽度方向)进行零填充(zero-padding),使其宽度统一为 max_width
    填充操作通过 torchvision.transforms.functional.pad 实现,其中填充参数 (0, pad_w, 0, 0) 表示:

    • 左侧填充 0,右侧填充 pad_w(即 max_width - 当前宽度
    • 上下方向不填充(均为 0)
      填充值为常数 0(即黑色背景),符合大多数 OCR 或图像识别任务的常规做法。
  4. 张量堆叠与标签处理
    经过上述处理后,所有图像具有相同的形状 [C, 32, max_width],因此可以使用 torch.stack(padding_images, dim=0) 沿 batch 维度堆叠成一个四维张量 images,形状为 [B, C, 32, max_width]
    同时,原始标签(通常为字符 ID 序列)被拼接为一维张量 targets = torch.cat(labels, dim=0),并额外返回每个样本标签的长度 target_lengths(类型为 LongTensor),以便后续在 CTC(Connectionist Temporal Classification)等损失函数中使用。

数据集推荐

数据集的格式一般是

-data

        -lable.txt     ###标签文件,一般格式是 "imgs/001.jpg 标签1"

        -imgs          ###图片文件夹

这里贴出一个200万条中英文图片的数据集路径,该数据集包含了100条英文文字识别数据和100万条中文文字识别数据

https://modelscope.cn/datasets/lzf010102/cheinese_and_englist_ocr_dataset_2000K

同时还有一个360万随机生成的中文训练数据集

文字识别_数据集-飞桨AI Studio星河社区

2.2 训练

为便于理解,我们将先对代码进行拆分讲解,完整的训练代码将在本章末尾统一贴出来。

main()

main方法主要是做一个搭建训练的整体流程,包括设备设置、数据加载、模型初始化、优化器与损失函数配置,并支持从检查点恢复训练;随后执行多轮训练与验证,根据验证损失保存最优模型,同时记录并可视化训练过程中的损失和准确率变化的功能。
 

def main():
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    train_dataloader, val_dataloader, chat_to_idx, idx_to_chat = load_data('merged_dataset/labels.txt',
                                                                           'merged_dataset/images')
    model = CRNNModel(in_channels=1, num_classes=len(chat_to_idx))
    model.to(device)

    checkpoint_path = 'best_crnn_model.pth'
    start_epoch = 0

    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4)
    scheduler = CosineAnnealingLR(optimizer, T_max=20, eta_min=1e-6)
    criterion = torch.nn.CTCLoss(blank=0, reduction='mean', zero_infinity=True)
    scaler = torch.cuda.amp.GradScaler()

    if os.path.exists(checkpoint_path):
        checkpoint = torch.load(checkpoint_path, map_location=device)
        if 'model_state_dict' in checkpoint:
            model.load_state_dict(checkpoint['model_state_dict'])
            print(f"Loaded model weights from {checkpoint_path}")
        else:
            model.load_state_dict(checkpoint)
            print(f"Loaded legacy model weights from {checkpoint_path}")
        if 'optimizer_state_dict' in checkpoint and 'epoch' in checkpoint:
            optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
            start_epoch = checkpoint['epoch'] + 1
            print("Restored optimizer state and epoch")

    epochs = 10
    best_val_loss = float('inf')

    # 初始化用于绘图的数据列表
    train_losses, val_losses, val_accuracies = [], [], []

    for epoch in range(start_epoch, epochs + 1):
        start_time = time.time()

        train_loss = train(model, device, train_dataloader, optimizer, criterion, epoch, scaler)
        train_losses.append(train_loss)

        val_loss, val_acc = validate(model, device, val_dataloader, criterion)
        val_losses.append(val_loss)
        val_accuracies.append(val_acc)

        scheduler.step()

        epoch_time = time.time() - start_time

        print(f"Epoch {epoch:02d}/{epochs} | "
              f"Time: {epoch_time:.2f}s | "
              f"Train Loss: {train_loss:.4f} | "
              f"Val Loss: {val_loss:.4f} | "
              f"Val Acc: {val_acc:.2%}")

        if val_loss < best_val_loss:
            best_val_loss = val_loss
            # 保存包括模型权重、优化器状态和当前epoch的完整检查点
            torch.save({
                'epoch': epoch,
                'model_state_dict': model.state_dict(),
                'optimizer_state_dict': optimizer.state_dict(),
            }, checkpoint_path)
            print(f"    → New best model saved! (Val Loss: {val_loss:.4f})")

        if epoch % 2 == 0:
            plot_training_history(train_losses, val_losses, val_accuracies)

    plot_training_history(train_losses, val_losses, val_accuracies)
1. 设备设置与数据加载
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
train_dataloader, val_dataloader, chat_to_idx, idx_to_chat = load_data(
    'merged_dataset/labels.txt', 'merged_dataset/images'
)
  • 首先判断当前环境是否支持 CUDA(即是否有可用的 GPU),若有则使用 GPU 加速训练,否则回退到 CPU。
  • 调用上文中的 load_data() 函数加载训练集和验证集的数据加载器(DataLoader),同时返回字符到索引(chat_to_idx)和索引到字符(idx_to_chat)的映射字典,用于后续解码模型输出。

2. 模型初始化
model = CRNNModel(in_channels=1, num_classes=len(chat_to_idx))
model.to(device)
  • 实例化我们的CRNN模型,这里的输入通道数为 1(因为是灰度图),分类类别数等于字表的大小。
  • 将模型移动到指定设备(GPU 或 CPU)上,确保后续计算在正确设备上执行。

3. 检查点路径与训练起始设置
checkpoint_path = 'best_crnn_model.pth'
start_epoch = 0
  • 定义模型检查点的保存路径。
  • 默认从第 0 轮(epoch = 0)开始训练,但如果存在检查点文件,则会从中恢复训练进度。

4. 训练组件配置
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4)
scheduler = CosineAnnealingLR(optimizer, T_max=20, eta_min=1e-6)
criterion = torch.nn.CTCLoss(blank=0, reduction='mean', zero_infinity=True)
scaler = torch.cuda.amp.GradScaler()
  • 优化器:使用 AdamW,相比传统 Adam 更适合带权重衰减的场景,有助于提升泛化能力。
  • 学习率调度器:采用 CosineAnnealingLR(余弦退火调度器),让学习率在训练过程中平滑下降,有助于模型收敛到更优解。
  • 损失函数:使用 CTCLoss(Connectionist Temporal Classification Loss),这是序列识别任务(如 OCR)的标准损失函数,能处理输入与输出长度不一致的问题。其中:
    • blank=0 表示空白符(CTC 中的特殊占位符)对应的索引;
    • zero_infinity=True 防止梯度爆炸(当 loss 为无穷大时自动置零)。
  • 混合精度训练:通过 GradScaler 启用自动混合精度(AMP),在保持模型精度的同时显著加速训练并减少显存占用(仅在 GPU 上有效)。

5. 从检查点恢复训练
if os.path.exists(checkpoint_path):
    checkpoint = torch.load(checkpoint_path, map_location=device)
    # 兼容新旧格式的检查点
    if 'model_state_dict' in checkpoint:
        model.load_state_dict(checkpoint['model_state_dict'])
        print(f"Loaded model weights from {checkpoint_path}")
    else:
        model.load_state_dict(checkpoint)  # 旧版仅保存模型权重
        print(f"Loaded legacy model weights from {checkpoint_path}")
    
    # 若包含优化器和 epoch 信息,则恢复完整训练状态
    if 'optimizer_state_dict' in checkpoint and 'epoch' in checkpoint:
        optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
        start_epoch = checkpoint['epoch'] + 1
        print("Restored optimizer state and epoch")
  • 如果存在检查点文件,程序会尝试加载模型权重、优化器状态和训练轮次,从而实现“断点续训”。
  • 同时兼容两种保存格式:仅模型权重(旧版)和完整训练状态(新版),提升代码鲁棒性。

6. 训练主循环
epochs = 10
best_val_loss = float('inf')
train_losses, val_losses, val_accuracies = [], [], []
  • 设置总训练轮数为 10。
  • 初始化 best_val_loss 为正无穷,用于跟踪验证损失的最低值。
  • 创建三个列表,用于记录每轮的训练损失、验证损失和验证准确率,便于后续绘图分析。

7. 每轮训练与验证
for epoch in range(start_epoch, epochs + 1):
    start_time = time.time()

    train_loss = train(model, device, train_dataloader, optimizer, criterion, epoch, scaler)
    train_losses.append(train_loss)

    val_loss, val_acc = validate(model, device, val_dataloader, criterion)
    val_losses.append(val_loss)
    val_accuracies.append(val_acc)

    scheduler.step()  # 更新学习率

    epoch_time = time.time() - start_time
    print(f"Epoch {epoch:02d}/{epochs} | Time: {epoch_time:.2f}s | "
          f"Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f} | Val Acc: {val_acc:.2%}")
  • 对每一轮:
    • 调用 train() 函数执行一次完整训练,返回平均训练损失;
    • 调用 validate() 函数在验证集上评估模型,返回验证损失和准确率;
    • 更新学习率(通过调度器);
    • 打印本轮训练的关键指标,便于实时监控。

8. 模型保存与可视化
if val_loss < best_val_loss:
    best_val_loss = val_loss
    torch.save({
        'epoch': epoch,
        'model_state_dict': model.state_dict(),
        'optimizer_state_dict': optimizer.state_dict(),
    }, checkpoint_path)
    print(f"    → New best model saved! (Val Loss: {val_loss:.4f})")

if epoch % 2 == 0:
    plot_training_history(train_losses, val_losses, val_accuracies)

# 最后再绘制一次完整曲线
plot_training_history(train_losses, val_losses, val_accuracies)
  • 模型保存策略:仅当当前验证损失低于历史最佳时,才保存模型。这种“早停+最优保存”策略能有效防止过拟合。
  • 可视化:每 2 个 epoch 绘制一次训练曲线,训练结束后再绘制最终完整曲线,帮助开发者直观分析模型收敛情况

ps:在训练的时候保存一定要保持完整的信息(模型权重,训练轮次、优化器状态等),因为博主在训练的过程中,云算力因为欠费而停机了一次,因为只保存了模型权重,重新训练的时候优化器等状态都没有,导致继续训练的时候新的第一轮loss直接起飞,从上一轮的0.5直接飙到0.8,虽然后续重新继续收敛,但是一直让博主很不爽

train()

train方法就是最主要的训练函数了,下面是train方法的代码

def train(model, device, train_loader, optimizer,criterion, epoch, scaler):
    model.train()
    total_loss = 0
    for index, (data, target, target_lengths) in enumerate(train_loader):
        data, target = data.to(device), target.to(device)
        optimizer.zero_grad()
        with torch.cuda.amp.autocast():
            output = model(data)
            ###ctc loss实现####
            output = output.log_softmax(2)
            input_length = torch.full((output.size(1),), output.size(0), dtype=torch.long)
            if target_lengths.max() > output.size(0):
                print("❗ Warning: Some target lengths exceed output time steps!")
                print(f"  Output time steps (T): {output.size(0)}")
                print(f"  Max target length in batch: {target_lengths.max().item()}")
                print(f"  Target lengths: {target_lengths.tolist()}")
                continue
            loss = criterion(output, target, input_length, target_lengths)
        ###反向传播###
        scaler.scale(loss).backward()
        scaler.step(optimizer)
        scaler.update()
        total_loss += loss.item()
        if index % 10 == 0:
            print(f'Train Epoch: {epoch} [{index * len(data)}/{len(train_loader.dataset)} '
                  f'({100. * index / len(train_loader):.0f}%)]\tLoss: {loss.item():.6f}')
    avg_loss = total_loss / len(train_loader)
    print(f'Train Epoch: {epoch} Average Loss: {avg_loss:.4f}')
    return avg_loss

下面是步骤解析:

1. 启用训练模式
model.train()
  • 调用 model.train() 将模型切换到训练模式。这会启用如 Dropout、BatchNorm 等层的训练行为,确保梯度正常计算。

2. 初始化损失累加器
total_loss = 0
  • 用于累计当前 epoch 中所有 batch 的损失值,以便最后计算平均训练损失。

3. 遍历训练数据
for index, (data, target, target_lengths) in enumerate(train_loader):
  • 使用 enumerate 遍历 train_loader,每次迭代返回一个 batch 的数据:
    • data:输入图像张量,形状通常为 (B, C, H, W)
    • target:目标标签序列,已转换为整数索引(如 [3, 5, 12, ...]),形状为 (sum(target_lengths),)(所有样本的标签拼接成一维);
    • target_lengths:每个样本的真实标签长度,形状为 (B,)

💡 注意:CTC 损失要求标签以“压缩一维”的形式输入(即所有样本的标签拼接),并配合 target_lengths 告知每个样本的原始长度。


4. 数据迁移到设备 & 梯度清零
data, target = data.to(device), target.to(device)
optimizer.zero_grad()
  • 将输入数据和标签移动到指定设备(GPU/CPU)。
  • 调用 optimizer.zero_grad() 清空上一轮的梯度缓存,防止梯度累积。

5. 前向传播
with torch.cuda.amp.autocast():
    output = model(data)
    output = output.log_softmax(2)
  • 使用 torch.cuda.amp.autocast() 上下文管理器启用自动混合精度训练
    • 在支持的运算中自动使用 float16(半精度)以节省显存、加速计算;
    • 关键操作(如 softmax、loss)仍保持数值稳定性。
  • 模型输出 output 的形状通常为 (T, B, C),其中:
    • T:时间步数(由 CNN + RNN 输出序列长度决定,这里就是经过压缩后的图片宽度);
    • B:batch size;
    • C:字符类别数(也就是字表的大小)。
  • 对输出在类别维度(dim=2)应用 log_softmax,这是 CTCLoss 的标准输入要求(需对数概率)。

6. 构造 CTC 损失所需的输入长度
input_length = torch.full((output.size(1),), output.size(0), dtype=torch.long)
  • input_length 表示每个样本在模型输出序列中的有效时间步数。
  • 由于 CRNN 通常对所有样本输出相同长度的序列(由图像宽度决定),因此每个样本的 input_length 都等于 T = output.size(0)
  • 使用 torch.full 创建一个长度为 B 的张量,所有元素值为 T

7. 安全检查:防止目标长度超过输出长度
if target_lengths.max() > output.size(0):
    print("❗ Warning: Some target lengths exceed output time steps!")
    # ... 打印调试信息并跳过该 batch
    continue
  • CTC 损失要求:对于每个样本,其真实标签长度不能超过模型输出的时间步数(即 target_length[i] ≤ T)。
  • 如果违反此条件,CTC 无法对齐,会报错或产生无效梯度。
  • 此处加入防御性检查:若发现异常 batch,打印警告信息并跳过该 batch,避免训练崩溃。

8. 计算 CTC 损失
loss = criterion(output, target, input_length, target_lengths)
  • 调用 CTCLoss 计算损失,参数说明:
    • output:模型输出的 log-probabilities,形状 (T, B, C)
    • target:拼接后的标签索引,形状 (N,)
    • input_length:每个样本的输出序列长度,形状 (B,)
    • target_lengths:每个样本的真实标签长度,形状 (B,)

具体的CTCLoss原理教程可以看

CTC Loss 数学原理讲解:Connectionist Temporal Classification-CSDN博客

这个大佬的讲解


9. 混合精度反向传播与优化
scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()
  • scaler.scale(loss).backward():对损失进行缩放后再反向传播,防止 float16 下梯度下溢(underflow)。
  • scaler.step(optimizer):在执行优化器更新前,自动将梯度反缩放回 float32 精度。
  • scaler.update():更新缩放因子,根据梯度是否溢出动态调整,确保训练稳定性。

🔧 这是使用 AMP 的标准三步操作,缺一不可。


10. 损失记录与进度打印
total_loss += loss.item()
if index % 10 == 0:
    print(f'Train Epoch: {epoch} [{index * len(data)}/{len(train_loader.dataset)} '
          f'({100. * index / len(train_loader):.0f}%)]\tLoss: {loss.item():.6f}')
  • 累加当前 batch 的损失(.item() 将张量转为 Python 数值,避免内存泄漏)。
  • 每 10 个 batch 打印一次训练进度和当前 loss,便于实时监控训练状态。

11. 返回平均损失
avg_loss = total_loss / len(train_loader)
print(f'Train Epoch: {epoch} Average Loss: {avg_loss:.4f}')
return avg_loss
  • 计算整个 epoch 的平均训练损失(总损失 / batch 数量)。
  • 打印并返回该值,供 main() 函数用于绘图或日志记录。
validate()

验证函数和训练函数差不多,只不过没有反向传播、更新梯度等更新模型参数的操作了,下面只贴出详细代码

def validate(model, device, val_loader, criterion):
    model.eval()
    total_loss = 0
    correct = 0
    total = 0
    with torch.no_grad():
        for index, (data, target, target_lengths) in enumerate(val_loader):
            data, target = data.to(device), target.to(device)

            output = model(data)
            output = output.log_softmax(2)

            input_length = torch.full((output.size(1),), output.size(0), dtype=torch.long)
            loss = criterion(output, target, input_length, target_lengths)
            total_loss += loss.item()

            # 解码预测结果(简单贪心解码)
            _, max_indices = torch.max(output, 2)
            predictions = []
            for seq in max_indices.permute(1, 0):
                # 移除重复字符和空白标签
                prev_char = None
                filtered_seq = []
                for char_idx in seq:
                    if char_idx != prev_char and char_idx != 0:
                        filtered_seq.append(char_idx.item())
                    prev_char = char_idx
                predictions.append(filtered_seq)

            # 计算准确率
            target = target.cpu().numpy()
            target_lengths = target_lengths.cpu().numpy()
            start = 0
            for i, pred in enumerate(predictions):
                length = target_lengths[i]
                true_seq = target[start:start + length]  # shape: (length,)
                start += length

                # 比较预测序列和真实序列是否完全一致
                if len(pred) == len(true_seq) and all(p == t for p, t in zip(pred, true_seq)):
                    correct += 1
                total += 1

    avg_loss = total_loss / len(val_loader)
    accuracy = 100. * correct / total if total > 0 else 0
    print(f'\nValidation set: Average loss: {avg_loss:.4f}, Accuracy: {correct}/{total} ({accuracy:.0f}%)\n')
    return avg_loss, accuracy

2.3 完整代码

至此,训练代码已全部完成,下面将贴出完整的训练代码包括数据处理类、工具文件、训练文件的代码

2.3.1 数据处理类
class CrnnDataSet(Dataset):
    def __init__(self, data_path, labels, chat_to_idx, transform=None):
        self.data_path = data_path
        self.labels = labels
        self.chat_to_idx = chat_to_idx
        self.transform = transform or transforms.Compose([
            transforms.Grayscale(),
            transforms.ToTensor(),
            transforms.Normalize((0.1,), (0.3,))
        ])

    def __len__(self):
        return len(self.data_path)

    def __getitem__(self, idx):
        img_path = self.data_path[idx]
        label = self.labels[idx]
        img = self.preprocess_image(img_path)
        filtered_label = []
        for c in label:
            if c in self.chat_to_idx:
                filtered_label.append(self.chat_to_idx[c])
        return {'image': img, 'label': torch.tensor(filtered_label, dtype=torch.long)}

    def preprocess_image(self,img_path):
        img = Image.open(img_path)
        img = self.transform(img)
        return img
2.3.2 工具文件
def load_data(file_path,file_root):
    img_data = []
    labels = []
    with open(file_path, 'r', encoding='c') as f:
        for line in f:
            img, label = line.strip().split(' ')
            img = os.path.join(file_root, img)
            img_data.append(img)
            labels.append(label)
    return img_data, labels



def init_chat_to_idx():
    chat_to_idx = {
        '<blank>': 0,
    }
    idx = len(chat_to_idx)
    with open('ppocr_keys_v1.txt', 'r', encoding='utf-8-sig') as f:
        for line in f:
            chat = line.strip()
            chat_to_idx[chat] = idx
            idx += 1
    idx_to_char = {v: k for k, v in chat_to_idx.items()}
    return chat_to_idx, idx_to_char
2.3.3 训练文件
import torch

import utils
from CrnnDataSet import CrnnDataSet
from torch.utils.data import DataLoader, Subset
from torch.nn import functional
import torchvision.transforms.functional as F
from sklearn.model_selection import train_test_split
from torch.nn.utils.rnn import pad_sequence
from torch import nn

from CRNNModel import CRNNModel
from torch.optim.lr_scheduler import CosineAnnealingLR,CosineAnnealingWarmRestarts
import matplotlib.pyplot as plt
import time
import os


def collate_fn(batches):
    images = []
    labels = []
    for batch in batches:
        images.append(batch['image'])
        labels.append(batch['label'])
    target_h = 32
    new_widths = []
    resized_images = []
    for img in images:
        c, h, w = img.shape
        scale = h / target_h
        new_width = int(w * scale)
        img_resized = F.resize(img, [target_h, new_width],antialias=True)
        resized_images.append(img_resized)
        new_widths.append(new_width)
    max_width = max(new_widths)
    padding_images = []
    for img in resized_images:
        pad_w = max_width - img.shape[-1]
        padding_img = functional.pad(img, (0, pad_w, 0, 0), 'constant', 0)
        padding_images.append(padding_img)
    images = torch.stack(padding_images, dim=0)
    targets = torch.cat(labels, dim=0)
    target_lengths = torch.LongTensor([len(label) for label in labels])
    return images, targets, target_lengths


def load_data(file_path, file_root):
    imgs, labels = utils.load_data(file_path, file_root)
    indices = list(range(len(imgs)))
    train_indices, val_indices = train_test_split(
        indices, test_size=0.01, random_state=42
    )
    chat_to_idx, idx_to_chat = utils.init_chat_to_idx()
    full_dataset = CrnnDataSet(imgs, labels, chat_to_idx)
    train_dataset = Subset(full_dataset, train_indices)
    val_dataset = Subset(full_dataset, val_indices)
    print(f"训练集大小: {len(train_dataset)}, 验证集大小: {len(val_dataset)}")
    train_dataloader = DataLoader(train_dataset, batch_size=128, shuffle=True, collate_fn=collate_fn, num_workers=4)
    val_dataloader = DataLoader(val_dataset, batch_size=128, shuffle=True, collate_fn=collate_fn, num_workers=4)
    return train_dataloader, val_dataloader, chat_to_idx, idx_to_chat

def train(model, device, train_loader, optimizer,criterion, epoch, scaler):
    model.train()
    total_loss = 0
    for index, (data, target, target_lengths) in enumerate(train_loader):
        data, target = data.to(device), target.to(device)
        optimizer.zero_grad()
        with torch.cuda.amp.autocast():
            output = model(data)
            ###ctc loss实现####
            output = output.log_softmax(2)
            input_length = torch.full((output.size(1),), output.size(0), dtype=torch.long)
            if target_lengths.max() > output.size(0):
                print("❗ Warning: Some target lengths exceed output time steps!")
                print(f"  Output time steps (T): {output.size(0)}")
                print(f"  Max target length in batch: {target_lengths.max().item()}")
                print(f"  Target lengths: {target_lengths.tolist()}")
                continue
            loss = criterion(output, target, input_length, target_lengths)
        ###反向传播###
        scaler.scale(loss).backward()
        scaler.step(optimizer)
        scaler.update()
        total_loss += loss.item()
        if index % 10 == 0:
            print(f'Train Epoch: {epoch} [{index * len(data)}/{len(train_loader.dataset)} '
                  f'({100. * index / len(train_loader):.0f}%)]\tLoss: {loss.item():.6f}')
    avg_loss = total_loss / len(train_loader)
    print(f'Train Epoch: {epoch} Average Loss: {avg_loss:.4f}')
    return avg_loss


def validate(model, device, val_loader, criterion):
    model.eval()
    total_loss = 0
    correct = 0
    total = 0
    with torch.no_grad():
        for index, (data, target, target_lengths) in enumerate(val_loader):
            data, target = data.to(device), target.to(device)

            output = model(data)
            output = output.log_softmax(2)

            input_length = torch.full((output.size(1),), output.size(0), dtype=torch.long)
            loss = criterion(output, target, input_length, target_lengths)
            total_loss += loss.item()

            # 解码预测结果(简单贪心解码)
            _, max_indices = torch.max(output, 2)
            predictions = []
            for seq in max_indices.permute(1, 0):
                # 移除重复字符和空白标签
                prev_char = None
                filtered_seq = []
                for char_idx in seq:
                    if char_idx != prev_char and char_idx != 0:
                        filtered_seq.append(char_idx.item())
                    prev_char = char_idx
                predictions.append(filtered_seq)

            # 计算准确率
            target = target.cpu().numpy()
            target_lengths = target_lengths.cpu().numpy()
            start = 0
            for i, pred in enumerate(predictions):
                length = target_lengths[i]
                true_seq = target[start:start + length]  # shape: (length,)
                start += length

                # 比较预测序列和真实序列是否完全一致
                if len(pred) == len(true_seq) and all(p == t for p, t in zip(pred, true_seq)):
                    correct += 1
                total += 1

    avg_loss = total_loss / len(val_loader)
    accuracy = 100. * correct / total if total > 0 else 0
    print(f'\nValidation set: Average loss: {avg_loss:.4f}, Accuracy: {correct}/{total} ({accuracy:.0f}%)\n')
    return avg_loss, accuracy




def main():
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    train_dataloader, val_dataloader, chat_to_idx, idx_to_chat = load_data('merged_dataset/labels.txt',
                                                                           'merged_dataset/images')
    model = CRNNModel(in_channels=1, num_classes=len(chat_to_idx))
    model.to(device)

    checkpoint_path = 'best_crnn_model.pth'
    start_epoch = 0

    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=1e-4)
    scheduler = CosineAnnealingLR(optimizer, T_max=20, eta_min=1e-6)
    criterion = torch.nn.CTCLoss(blank=0, reduction='mean', zero_infinity=True)
    scaler = torch.cuda.amp.GradScaler()

    if os.path.exists(checkpoint_path):
        checkpoint = torch.load(checkpoint_path, map_location=device)
        if 'model_state_dict' in checkpoint:
            model.load_state_dict(checkpoint['model_state_dict'])
            print(f"Loaded model weights from {checkpoint_path}")
        else:
            model.load_state_dict(checkpoint)
            print(f"Loaded legacy model weights from {checkpoint_path}")
        if 'optimizer_state_dict' in checkpoint and 'epoch' in checkpoint:
            optimizer.load_state_dict(checkpoint['optimizer_state_dict'])
            start_epoch = checkpoint['epoch'] + 1
            print("Restored optimizer state and epoch")

    epochs = 10
    best_val_loss = float('inf')

    # 初始化用于绘图的数据列表
    train_losses, val_losses, val_accuracies = [], [], []

    for epoch in range(start_epoch, epochs + 1):
        start_time = time.time()

        train_loss = train(model, device, train_dataloader, optimizer, criterion, epoch, scaler)
        train_losses.append(train_loss)

        val_loss, val_acc = validate(model, device, val_dataloader, criterion)
        val_losses.append(val_loss)
        val_accuracies.append(val_acc)

        scheduler.step()

        epoch_time = time.time() - start_time

        print(f"Epoch {epoch:02d}/{epochs} | "
              f"Time: {epoch_time:.2f}s | "
              f"Train Loss: {train_loss:.4f} | "
              f"Val Loss: {val_loss:.4f} | "
              f"Val Acc: {val_acc:.2%}")

        if val_loss < best_val_loss:
            best_val_loss = val_loss
            # 保存包括模型权重、优化器状态和当前epoch的完整检查点
            torch.save({
                'epoch': epoch,
                'model_state_dict': model.state_dict(),
                'optimizer_state_dict': optimizer.state_dict(),
            }, checkpoint_path)
            print(f"    → New best model saved! (Val Loss: {val_loss:.4f})")

        if epoch % 2 == 0:
            plot_training_history(train_losses, val_losses, val_accuracies)

    plot_training_history(train_losses, val_losses, val_accuracies)


def plot_training_history(train_losses, val_losses, val_accuracies):
    epochs = range(1, len(train_losses) + 1)

    plt.figure(figsize=(12, 4))

    # 损失曲线
    plt.subplot(1, 2, 1)
    plt.plot(epochs, train_losses, label='Train Loss', marker='o')
    plt.plot(epochs, val_losses, label='Validation Loss', marker='o')
    plt.title('Loss over Epochs')
    plt.xlabel('Epoch')
    plt.ylabel('Loss')
    plt.legend()
    plt.grid(True)

    # 准确率曲线
    plt.subplot(1, 2, 2)
    plt.plot(epochs, val_accuracies, label='Validation Accuracy', color='green', marker='o')
    plt.title('Validation Accuracy over Epochs')
    plt.xlabel('Epoch')
    plt.ylabel('Accuracy')
    plt.ylim(0, 1)
    plt.legend()
    plt.grid(True)

    plt.tight_layout()
    plt.savefig('training_history.png')
    plt.show()
    print("\nTraining history plot saved as 'training_history.png'")


if __name__ == '__main__':
    main()


三、推理

在上面两章中,我们完成了模型网络的搭建以及模型的训练,接下来我们该试一下自己训练的模型有什么效果

推理部分有两个要注意的地方

一个是对图片的处理,下面看代码

 def preprocess_image(self, image_path, max_height=32, max_width=128):
        """预处理方法"""
        img = Image.open(image_path)
        width, height = img.size
        scale_factor = max_height / height
        new_width = int(width * scale_factor)
        img = img.resize((new_width, max_height))
        img = self.transform(img)
        return img

是不是很熟悉,没错,推理时候的处理要和训练时对图片的处理一样,这是因为模型是在特定预处理方式下训练的,推理时若处理方式不同,会导致输入分布偏移,模型无法正确识别,甚至出错或崩溃。所以必须保持一致。

另一个则是解码模型输出的解码器,代码是:

    def decode_predictions(self, output):
        """解码模型输出为文本"""
        _, max_indices = torch.max(output, 2)  # output形状: (T, N, C)
        max_indices = max_indices.permute(1, 0)  # 转换为(N, T)

        predictions = []
        for seq in max_indices:
            prev_char = None
            filtered_seq = []
            for char_idx in seq:
                char_idx = char_idx.item()
                if char_idx != prev_char and char_idx != 0:  # 0是空白标签
                    filtered_seq.append(char_idx)
                prev_char = char_idx
            prediction = ''.join([self.idx_to_chat[idx] for idx in filtered_seq])
            predictions.append(prediction)

        return predictions

这里使用了贪心解码:

  1. 对每个时间步,选择概率最高的字符(即 argmax)。
  2. 合并连续重复的非空白字符(因为 CTC 允许通过重复表示同一个字符,如 "hhheeellllo" → "hello")。
  3. 移除所有空白标签(索引为 0 的字符)。

下面是推理文件的完整代码:

import os
import torch
from PIL import Image
from torchvision import transforms
from CRNNModel import CRNNModel
from utils import init_chat_to_idx


class CRNNInference:
    def __init__(self, model_path, device='cuda' if torch.cuda.is_available() else 'cpu'):
        self.device = device
        self.chat_to_idx, self.idx_to_chat = init_chat_to_idx()
        print("Blank label index:", self.chat_to_idx.get('', -1))
        self.num_classes = len(self.chat_to_idx)

        # 初始化模型
        self.model = CRNNModel(in_channels=1, num_classes=self.num_classes)
        checkpoint = torch.load(model_path, map_location=device)

        # 判断checkpoint格式并加载相应的数据
        if 'model_state_dict' in checkpoint:
            self.model.load_state_dict(checkpoint['model_state_dict'])
        else:
            self.model.load_state_dict(checkpoint)
        self.model.to(device)
        self.model.eval()

        # 图像预处理
        self.transform = transforms.Compose([
            transforms.Grayscale(),
            transforms.ToTensor(),
            transforms.Normalize((0.1,), (0.3,))
        ])

    def preprocess_image(self, image_path, max_height=32, max_width=128):
        """预处理方法"""
        img = Image.open(image_path)
        width, height = img.size
        scale_factor = max_height / height
        new_width = int(width * scale_factor)
        img = img.resize((new_width, max_height))
        img = self.transform(img)
        return img

    def decode_predictions(self, output):
        """解码模型输出为文本"""
        _, max_indices = torch.max(output, 2)  # output形状: (T, N, C)
        max_indices = max_indices.permute(1, 0)  # 转换为(N, T)

        predictions = []
        for seq in max_indices:
            prev_char = None
            filtered_seq = []
            for char_idx in seq:
                char_idx = char_idx.item()
                if char_idx != prev_char and char_idx != 0:  # 0是空白标签
                    filtered_seq.append(char_idx)
                prev_char = char_idx
            prediction = ''.join([self.idx_to_chat[idx] for idx in filtered_seq])
            predictions.append(prediction)

        return predictions

    def predict(self, image_path):
        """对单个图像进行预测"""
        image_tensor = self.preprocess_image(image_path).to(self.device)
        print(image_tensor.shape)
        image_tensor = image_tensor.unsqueeze(0)

        with torch.no_grad():
            output = self.model(image_tensor)
            print(output.shape)
            output = output.log_softmax(2)

        predictions = self.decode_predictions(output)
        return predictions[0] if predictions else ""


if __name__ == '__main__':
    model_path = ('best_crnn_model.pth')
    crnn = CRNNInference(model_path)

    test_folder = 'test_data'
    image_paths = [os.path.join(test_folder, f) for f in os.listdir(test_folder) if
                   os.path.isfile(os.path.join(test_folder, f))]
    image_paths.sort()
    print(image_paths)

    for image_path in image_paths:
        try:
            prediction = crnn.predict(image_path)
            print(f"Prediction for {image_path}: {prediction}")
        except Exception as e:
            print(f"Could not process file {image_path}: {e}")

四、总结

到这里,本文章已经全部完成了,这里分享一些博主训练的数据

显卡使用的是autodl的5090,在200万个图片的数据集上训练了28个epoch达到收敛,这时训练集的loss在0.12左右,验证集的loss在0.15左右,验证集的正确率达到了89%左右,共训练了27个小时(全精度)

可以看看效果:

1:

2:

3:长文本识别效果

4.英文场景

正确率还是很高的,感觉上比验证集的89%要高得多

但是还是有不足:由于训练集中没有手写字符、中英文混合图片,导致模型在非常抽象的手写字符下以及中英文混合场景下的识别效果比较差

最后,如果感觉文章中的代码分布很乱,找不到哪些代码应该放在那里,下面贴出已经上传到modelscope的模型文件和上面演示的权重

中英文文字识别模型 · 模型库https://modelscope.cn/models/lzf010102/lzf_ocr_v1如果有错误或者有更好的优化方法,欢迎大佬指证

后续学习想法:本来想基于qwen3的模型网络架构去从头预训练+微调训练一个llm,但是跑起来后发现仅仅400万条文本数据训练5个epoch就得用5090训练70多个小时,由于囊中羞涩,所以这个计划暂时搁置了🙄🙄🙄

后续可能还会基于transformers去实现一个视觉模型😐😐

更多推荐