news 2026/8/28 10:11:55

从零部署OCR系统:EAST+CRNN端到端文本检测与识别实战

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
从零部署OCR系统:EAST+CRNN端到端文本检测与识别实战

简介:OCR(光学字符识别)技术旨在将图像中的文字信息转换为可编辑的文本数据,其核心原理是通过计算机视觉和深度学习模型模拟人类的阅读过程。该技术通过特征提取、序列建模和解码等步骤,实现了对复杂场景下文本的自动化理解,具有极高的工程应用价值。在文档数字化、车牌识别、工业自动化等场景中,OCR技术能显著提升信息处理效率与准确性。本文聚焦于结合EAST文本检测模型与CRNN+CTC文本识别模型的经典解决方案,详细解析了从环境搭建、模型原理到推理Pipeline构建的完整流程,并针对TensorFlow版本兼容性、模型加载等常见陷阱提供了实战经验,为开发者构建本地化、高定制化的OCR系统提供了清晰的路径。

1. 从零部署一套完整的OCR系统:为什么选择EAST+CRNN?

最近在做一个需要批量处理图片中文字信息的项目,从网上找了不少开源方案,发现很多教程要么只讲检测,要么只讲识别,要么环境配置写得云里雾里。折腾了好几天,终于把基于Keras和TensorFlow的EAST/AdvancedEAST文本检测模型和CRNN+CTC文本识别模型这套组合拳给跑通了。今天就把从环境搭建到模型推理的完整流程,以及我踩过的那些坑,详细梳理一遍。

这套方案的核心价值在于,它提供了一个从“找到文字在哪里”到“认出文字是什么”的端到端解决方案。EAST(Efficient and Accurate Scene Text detector)负责在复杂背景的图片中精准定位文字区域,无论是水平文本还是有一定倾斜角度的文本,它都能用四边形(或旋转矩形)框出来。而CRNN(Convolutional Recurrent Neural Network)结合CTC(Connectionist Temporal Classification)损失函数,则擅长处理不定长的文本序列识别,特别适合识别从图片中裁剪出来的单个文本行。

你可能会问,现在不是有很多现成的OCR接口吗?没错,但对于需要本地化部署、处理敏感数据、或者对识别精度和速度有定制化需求的场景,自己搭建一套可控的模型 pipeline 是非常有必要的。而且,通过理解这套经典组合的内部机制,你能更好地应对各种“奇葩”的图片,比如低光照、模糊、艺术字体等,这是调用黑盒API所无法获得的灵活性。

2. 环境搭建:避开TensorFlow 2.x的版本陷阱

万事开头难,环境配置是第一个拦路虎。项目基于Keras和TensorFlow,但这两个库的版本兼容性问题堪称经典。直接pip install tensorflow装最新版(比如2.18),大概率会跑不起来,因为原始代码往往是为老版本(如TensorFlow 1.x或特定的2.x早期版本)编写的。

我的经验是,优先使用虚拟环境,这能保证你的项目依赖与系统其他Python环境隔离。这里我用的是conda,用venv也一样。

# 创建一个新的虚拟环境,指定Python版本为3.7(兼容性较好) conda create -n ocr_env python=3.7 conda activate ocr_env

接下来安装TensorFlow和Keras。经过多次测试,一个比较稳定的组合是TensorFlow 2.3.0 和 Keras 2.4.3。这个组合既保留了2.x API的主要特性,又对许多老代码有较好的兼容性。

pip install tensorflow==2.3.0 pip install keras==2.4.3

注意:这里安装的keras包实际上是tf.keras的一个独立封装。在代码中,我们通常会直接使用from tensorflow import keras,以确保使用的是TensorFlow内部的Keras实现,避免冲突。

为什么是2.3.0?在TensorFlow 2.x的演进中,2.0到2.2版本变动剧烈,很多API不稳定。2.3版本是一个相对成熟的节点,对tf.compat.v1(兼容1.x代码的模块)的支持也比较完善。而Keras 2.4.3是与TF 2.3匹配的版本。盲目安装最新的2.18,你可能会遇到诸如tf.placeholdertf.Session等1.x的符号找不到,或者Keras层与TF不兼容等各种错误,排查起来极其耗时。

安装完核心框架后,还需要一些辅助库:

pip install opencv-python-headless # 图像处理,用headless版本避免GUI依赖 pip install numpy pip install Pillow # 图像处理 pip install shapely # 用于处理几何图形(如多边形框) pip install pyclipper # 依赖Shapely,用于框的缩放计算

如果遇到pyclipper安装失败,可能是缺少C++编译环境。在Windows上,可以尝试安装Microsoft Visual C++ Build Tools;在Linux/macOS上,确保安装了gccpython3-dev

3. EAST与AdvancedEAST文本检测模型深度解析

环境搞定,我们来深入看看文本检测部分。EAST模型之所以在当年引起关注,是因为它摒弃了当时主流的“多步骤”检测框架(如CTPN),转而采用一种“端到端”的、全卷积网络(FCN)的架构,速度更快,结构更优雅。

3.1 EAST模型的核心思想:直接回归几何框

传统的文本检测方法,可能需要先产生候选框,再分类和回归修正。EAST的思路很直接:让网络在图像的每个像素点(更准确地说,是特征图上的每个点)上,直接预测两个东西:

  1. 文本得分:该像素点是否位于文本区域内的概率。
  2. 几何形状:如果该点是文本点,那么它到包含该点的文本外接矩形的四条边的距离(对于旋转框,则是到四边的距离加上一个旋转角度)。

在推理时,模型会输出一个密集的预测图。然后通过一个简单的后处理步骤(称为“Locality-Aware NMS”,局部感知非极大值抑制),将这些密集的、重叠的预测框合并成最终的、稀疏的文本检测框。

网络结构通常基于一个强大的卷积主干网络(如PVANet或,在开源实现中常用的是VGG16或ResNet的前几层)来提取特征。然后通过一个特征金字塔网络(FPN)的思想,将深层语义特征和浅层位置特征融合,最终通过几个卷积层输出我们想要的文本得分和几何形状通道。

3.2 AdvancedEAST的改进:应对长文本挑战

原始的EAST模型在处理长文本行时,有时会将其断裂成多个小段。AdvancedEAST(也称AEAST)主要针对这个问题进行了改进。

它的核心改动在于几何形状的表示方式。EAST预测每个点到文本框四边的距离((d_top, d_right, d_bottom, d_left)),这被称为“RBOX”(旋转框)表示法。而AdvancedEAST增加了一种“QUAD”(四边形)的表示法,直接预测文本框四个顶点的坐标偏移。对于形状不规则或特别长的文本,四边形的表示法更加灵活。

在实际代码中,AdvancedEAST的网络会同时输出RBOX和QUAD两种几何信息,在训练时根据标注数据选择合适的损失进行计算。在推理时,可以优先使用QUAD的结果,以获得对长文本更好的包围效果。

3.3 模型文件与加载

你下载的EAST_AdvancedEAST压缩包里,应该包含.h5.pb格式的模型权重文件。在Keras中,加载.h5文件非常简单:

from tensorflow.keras.models import load_model # 加载EAST模型 east_model = load_model('east_model.h5', compile=False) # compile=False可以避免加载优化器状态,加快速度 # 或者加载AdvancedEAST模型 advanced_east_model = load_model('advanced_east_model.h5', compile=False)

这里有个关键点:compile=False。因为我们只是用模型进行推理(预测),不需要训练时的优化器、损失函数等配置。加上这个参数能避免一些不必要的警告,并且加载速度更快。

如果模型文件是TensorFlow的SavedModel格式(一个包含saved_model.pb和变量文件夹的目录),则加载方式不同:

import tensorflow as tf model = tf.saved_model.load('path_to_saved_model_directory') # 调用时使用 model.signatures['serving_default']

你需要根据压缩包内的实际文件结构来决定加载方式。通常,.h5文件更为常见。

4. CRNN+CTC文本识别模型原理与实现

检测模型把文字框抠出来了,下一步就是识别框里的内容。这就是CRNN+CTC的舞台。

4.1 CRNN的网络结构:卷积、循环、转录的三重奏

CRNN,顾名思义,由三部分组成:

  1. 卷积层(CNN):使用深度卷积网络(如VGG的变种)从输入图像(裁剪出的文本行图像)中提取视觉特征序列。输入图像被归一化到相同高度(如32像素),宽度可变。CNN输出的特征图在高度上被池化到1,宽度维度则对应了输入图像从左到右的序列。例如,一个100像素宽的图,经过CNN后可能得到一个[1, 25, 512]的特征张量(1是高度,25是序列长度,512是特征通道数)。这相当于把图像转换成了一串25个“特征向量”,每个向量代表了图像水平方向上一小列区域的视觉信息。
  2. 循环层(RNN):将上一步得到的特征序列(25个512维向量)输入到双向循环神经网络(常用LSTM)中。RNN的优势在于处理序列数据,它能够结合上下文信息,对当前时刻的特征进行编码。双向LSTM会同时考虑从左到右和从右到左的上下文,这对于识别字符非常关键,因为字符的识别往往依赖于其相邻字符(比如“i”和“l”的区分)。
  3. 转录层(Transcription):将RNN输出的序列,映射到最终的字符序列。这里就是CTC大显身手的地方。

4.2 CTC损失函数:解决序列对齐的魔法

这是整个识别模型最精妙也最难理解的部分。简单来说,它让网络无需事先知道输入特征序列和输出标签序列之间精确的逐帧对齐关系

网络在RNN之后,通常会接一个全连接层,输出每个时间步上所有可能字符(包括一个特殊的“空白”标签-)的概率分布。假设字符表有68个字符(62个字母数字+5个标点+1个空白),序列长度是25,那么输出形状就是[25, 68]

直接取每个时间步概率最大的字符,会得到一串25个字符的序列,里面可能包含大量重复字符和空白符,例如--hh--e--lll---lo--。CTC解码的任务就是,将这样的序列“压缩”成最终的标签序列。规则是:

  • 移除重复的字符(除非它们之间有空白符隔开)。
  • 然后移除所有的空白符-

按照这个规则,--hh--e--lll---lo--就会被压缩成hello。CTC损失函数在训练时,就是计算网络输出的概率分布,能通过这种规则得到正确标签hello的总概率,并最大化这个概率。它自动处理了对齐问题,我们只需要提供图像和对应的标签文本即可,无需标注出每个字符在图像中的具体位置。

4.3 加载与使用CRNN模型

CRNN模型同样通常以.h5格式提供。加载方式与检测模型类似:

crnn_model = load_model('crnn_model.h5', compile=False)

但是,这里有一个巨大的坑:很多开源的CRNN模型在保存时,包含了CTC损失层。而Keras/TensorFlow的CTC损失函数(ctc_batch_costLambda层封装的CTC)在保存和加载时容易出现问题,尤其是在compile=False的情况下,模型输出可能会不符合预期。

稳妥的做法是,重新定义模型结构,然后只加载权重。假设你知道原始模型的结构定义(这通常能在源代码中找到):

from tensorflow.keras.models import Model from tensorflow.keras.layers import Input, Conv2D, MaxPooling2D, Reshape, Bidirectional, LSTM, Dense, Lambda import tensorflow as tf # 1. 定义模型结构(此处为示例,需根据实际模型调整) def build_crnn_model(input_shape, num_classes): input_img = Input(shape=input_shape, name='input_image') # ... 这里省略具体的CNN和RNN层构建过程 ... # 假设最终得到输出层,名为 ‘output’ # output = Dense(num_classes, activation='softmax')(...) # 注意:原始的用于训练的模型可能包含Lambda(ctc_loss)层,用于推理的模型不应该包含这一层。 # 我们构建的应该是“推理模型”,它直接输出RNN后的特征序列或经过softmax的概率。 # 示例:构建一个仅输出特征序列的模型(需与权重文件匹配) model = Model(inputs=input_img, outputs=output, name='CRNN_Inference') return model # 2. 实例化模型 height, width, channel = 32, 100, 1 # 假设输入图像高32,宽100,通道1(灰度) num_classes = 68 # 字符表大小 inference_model = build_crnn_model((height, width, channel), num_classes) # 3. 加载权重 inference_model.load_weights('crnn_model_weights.h5') # 注意这里加载的是权重文件,不是整个模型

更常见的情况是,作者会提供两个模型文件:一个包含CTC损失层的用于训练(crnn_train.h5),一个剥离了损失层、输出更干净的直接用于推理(crnn_inference.h5)。务必使用推理模型。

如果只有单个.h5文件,加载后调用model.predict发现输出维度很奇怪(比如多出一个维度),很可能加载的是训练模型。你需要查看模型结构model.summary(),并可能手动提取出内部的子模型作为推理模型。

5. 实战演练:构建端到端OCR推理Pipeline

理论说得再多,不如跑通代码。下面我们一步步搭建一个完整的OCR流程。

5.1 步骤一:使用EAST模型检测文本区域

首先,我们需要对输入图片进行预处理,使其符合EAST模型的输入要求,然后进行预测和后处理。

import cv2 import numpy as np from tensorflow.keras.preprocessing.image import img_to_array def preprocess_for_east(image, max_side_len=2400): """ 预处理图像供EAST模型使用。 1. 将图像等比例缩放,长边不超过max_side_len。 2. 填充图像至32的倍数(网络下采样倍数)。 """ h, w = image.shape[:2] resize_w = w resize_h = h # 限制长边 ratio = max_side_len / max(h, w) if ratio < 1: resize_h = int(h * ratio) resize_w = int(w * ratio) # 确保新尺寸是32的倍数 resize_h = resize_h if resize_h % 32 == 0 else (resize_h // 32) * 32 resize_w = resize_w if resize_w % 32 == 0 else (resize_w // 32) * 32 image_resized = cv2.resize(image, (resize_w, resize_h)) # 归一化 image_resized = image_resized.astype(np.float32) image_resized -= np.array([103.939, 116.779, 123.68]) # ImageNet均值 # 如果模型训练时输入是RGB,而OpenCV读入是BGR,需要转换 # image_resized = image_resized[..., ::-1] # BGR to RGB # 增加批次维度 image_resized = np.expand_dims(image_resized, axis=0) return image_resized, (h / resize_h, w / resize_w) # 返回缩放比例 def decode_east_predictions(score_map, geo_map, score_thresh=0.8, nms_thresh=0.2): """ 解码EAST模型的预测输出。 score_map: 文本得分图 [H, W, 1] geo_map: 几何信息图 [H, W, 5] (5: d_top, d_right, d_bottom, d_left, angle) 返回:检测框列表,每个框为 [x1, y1, x2, y2, x3, y3, x4, y4] (四点坐标) """ # 1. 根据得分阈值筛选像素点 xy_text = np.argwhere(score_map > score_thresh) # [n, 2] 格式为 [y, x] if len(xy_text) == 0: return [] # 2. 根据这些点的几何信息还原文本框 # 这里涉及复杂的几何计算,包括根据角度旋转等。 # 开源实现中通常有现成的函数,例如从EAST官方代码或流行复现中移植。 # 以下为简化伪代码逻辑: boxes = [] for y, x in xy_text: d_top, d_right, d_bottom, d_left, angle = geo_map[y, x] # 根据点到四边的距离和角度,计算原始图像中对应的四边形顶点... # ... # boxes.append([x1, y1, x2, y2, x3, y3, x4, y4]) pass # 3. 应用Locality-Aware NMS合并重叠框 # 这也是一个标准步骤,有现成实现。 # final_boxes = lanms.merge_quadrangle_n9(boxes, nms_thresh) return final_boxes

由于解码和NMS的实现较为复杂且固定,强烈建议直接使用可靠的第三方实现,例如在GitHub上搜索EAST text detection decode找到的相关代码。你的模型压缩包内很可能也包含了这些工具函数。

5.2 步骤二:文本区域裁剪与矫正

EAST输出的框可能是旋转的四边形。为了给CRNN识别,我们需要将这些区域“拉直”成水平的矩形图像。

def four_points_transform(image, pts): """ 透视变换,将四边形区域矫正为矩形。 pts: 形状为(4, 2)的np数组,四个顶点坐标。 返回:矫正后的矩形图像。 """ # 将顶点排序为:左上,右上,右下,左下 rect = order_points(pts) (tl, tr, br, bl) = rect # 计算新矩形的宽度和高度 widthA = np.sqrt(((br[0] - bl[0]) ** 2) + ((br[1] - bl[1]) ** 2)) widthB = np.sqrt(((tr[0] - tl[0]) ** 2) + ((tr[1] - tl[1]) ** 2)) maxWidth = max(int(widthA), int(widthB)) heightA = np.sqrt(((tr[0] - br[0]) ** 2) + ((tr[1] - br[1]) ** 2)) heightB = np.sqrt(((tl[0] - bl[0]) ** 2) + ((tl[1] - bl[1]) ** 2)) maxHeight = max(int(heightA), int(heightB)) # 目标点坐标 dst = np.array([ [0, 0], [maxWidth - 1, 0], [maxWidth - 1, maxHeight - 1], [0, maxHeight - 1]], dtype="float32") # 计算透视变换矩阵并应用 M = cv2.getPerspectiveTransform(rect, dst) warped = cv2.warpPerspective(image, M, (maxWidth, maxHeight)) return warped def order_points(pts): """ 将四个点排序为:左上,右上,右下,左下。 """ rect = np.zeros((4, 2), dtype="float32") s = pts.sum(axis=1) rect[0] = pts[np.argmin(s)] # 左上角点x+y最小 rect[2] = pts[np.argmax(s)] # 右下角点x+y最大 diff = np.diff(pts, axis=1) rect[1] = pts[np.argmin(diff)] # 右上角点x-y最小 rect[3] = pts[np.argmax(diff)] # 左下角点x-y最大 return rect

对于每个检测到的四边形框,使用four_points_transform函数即可得到矫正后的文本行图像。

5.3 步骤三:CRNN识别与CTC解码

将矫正后的文本行图像预处理后,送入CRNN模型进行识别。

def preprocess_for_crnn(image, target_height=32): """ 预处理单张文本行图像供CRNN使用。 1. 转换为灰度图。 2. 调整高度至target_height,宽度按比例缩放。 3. 归一化并转置维度为 [H, W, C] -> [W, H, C]? (取决于模型输入) 注意:CRNN模型的输入格式需要根据训练时的设置来确定,常见的是 [H, W, C] 且宽度可变。 """ if len(image.shape) == 3: image = cv2.cvtColor(image, cv2.COLOR_BGR2GRAY) h, w = image.shape ratio = target_height / h new_w = int(w * ratio) resized = cv2.resize(image, (new_w, target_height), interpolation=cv2.INTER_CUBIC) # 归一化到[-1, 1]或[0, 1] resized = resized.astype(np.float32) / 255.0 # 有些模型要求输入均值为0.5,标准差为0.5 # resized = (resized - 0.5) / 0.5 # 增加通道维和批次维 resized = np.expand_dims(resized, axis=-1) # [H, W, 1] resized = np.expand_dims(resized, axis=0) # [1, H, W, 1] # 注意:有些CRNN实现要求输入为 [batch, height, width, channel] # 有些则要求 [batch, width, height, channel] 或 [batch, channel, height, width],务必与模型对齐! return resized, new_w def decode_ctc_predictions(preds, charset): """ 解码CRNN模型的CTC输出。 preds: 模型输出的概率矩阵,形状为 [序列长度, 字符类别数+1(空白符)] charset: 字符列表,索引与模型输出类别对应。 返回:识别出的字符串。 """ # 方法1:贪婪解码(取每个时间步最大概率的字符) pred_indices = np.argmax(preds, axis=-1) # [序列长度,] # 方法2:使用CTC束搜索解码(更准确,但稍慢) # 这里演示贪婪解码 last_char = -1 text = [] for idx in pred_indices: if idx != len(charset): # 如果不是空白符 if idx != last_char: # 去除连续重复字符 text.append(charset[idx]) last_char = idx else: last_char = -1 # 遇到空白符,重置last_char return ''.join(text) # 假设我们已经有了推理模型 crnn_inference_model 和字符表 char_list def recognize_text(image_cropped): processed_img, width = preprocess_for_crnn(image_cropped, target_height=32) # 确保输入维度与模型匹配,这里假设输入为 [batch, height, width, channel] # 如果模型输入是 [batch, width, height, channel],需要转置 # processed_img = np.transpose(processed_img, (0, 2, 1, 3)) preds = crnn_inference_model.predict(processed_img, verbose=0) # preds 形状可能是 [1, 序列长度, 字符数],需要 squeeze 掉批次维 preds = np.squeeze(preds, axis=0) text = decode_ctc_predictions(preds, char_list) return text

字符表char_list必须与模型训练时使用的完全一致,通常包含数字、大小写字母和常见标点符号。这个信息一般会在模型提供的说明或源代码中定义。

6. 性能优化与常见问题排查

将检测和识别串联起来后,一个基础的OCR系统就完成了。但在实际使用中,你肯定会遇到性能和精度问题。

6.1 检测阶段优化

  • 多尺度检测:EAST模型对尺度敏感。对于一张图里文字大小差异很大的情况,可以对原图构建图像金字塔(缩放到多个不同尺寸),分别检测后再将结果映射回原图坐标,最后合并。这能显著提升大小文字的共同检出率,但计算量会成倍增加。
  • 调整得分阈值score_thresh控制检测框的置信度。调高它会减少误检(把非文字区域框出来),但可能漏检一些模糊文字;调低它则相反。需要根据你的场景在精确率(Precision)和召回率(Recall)之间权衡。
  • NMS阈值调整nms_thresh控制框合并的激进程度。值越小,合并越严格,同一个文字区域只保留一个框;值越大,越可能保留多个重叠框。对于文字间距很小的场景,可以适当调大。

6.2 识别阶段优化

  • 输入图像高度:CRNN模型通常固定了输入图像的高度(如32像素)。在预处理时,必须严格按照这个高度进行等比例缩放。缩放算法建议使用cv2.INTER_CUBIC,它在放大时效果较好。
  • 图像二值化/增强:对于低质量文本图像,在送入CRNN前可以进行一些预处理,如自适应二值化、对比度拉伸、去噪等,能有效提升识别率。OpenCV的cv2.createCLAHE(对比度受限自适应直方图均衡化)对光照不均的图片很有效。
  • 字符表完整性:确保你的char_list包含了所有可能出现的字符。如果遇到模型不认识的字符,它可能会识别成乱码或空白。对于中文场景,需要训练包含汉字的CRNN模型,字符表会非常大(几千字)。

6.3 经典错误与解决方案

  1. 模型加载失败,提示Unknown layer: LambdaCTC

    • 原因:模型保存时包含了自定义层或Lambda层(如CTC损失层),加载时找不到定义。
    • 解决:在load_model时使用custom_objects参数传入自定义层。或者,更推荐的方法是按前面所述,重建推理模型结构并只加载权重。
    # 如果必须加载完整模型,且知道自定义层 from tensorflow.keras.layers import Lambda import tensorflow as tf def ctc_lambda_func(args): # 这里定义与保存模型时一致的CTC Lambda层逻辑 y_pred, labels, input_length, label_length = args return tf.keras.backend.ctc_batch_cost(labels, y_pred, input_length, label_length) custom_objects = {'ctc_lambda_func': ctc_lambda_func, 'Lambda': Lambda} model = load_model('crnn.h5', custom_objects=custom_objects, compile=False)
  2. CRNN识别结果全是乱码或重复字符

    • 原因a:输入图像预处理与模型训练时不匹配。可能是归一化方式(/255.0还是(img-mean)/std)、通道顺序(RGB vs BGR)、图像宽度缩放算法不对。
    • 解决:仔细检查模型训练代码的预处理部分,并完全复现。一个常见的错误是,训练时用了PIL(RGB)读图,推理时用了OpenCV(BGR)但没转换。
    • 原因b:字符表顺序不对。模型输出第0类对应什么字符,必须与解码时char_list的第0个元素一致。
    • 解决:找到模型训练时生成char_list的代码,确保完全一致。
  3. EAST检测框歪斜或包含非文字区域

    • 原因:后处理参数(score_thresh,nms_thresh)或解码过程中的几何计算参数(如文本框最小面积、边长比限制)设置不合理。
    • 解决:可视化中间结果。将score_map以热力图形式显示出来,看看高亮区域是否对应文字。调整上述参数,并在你的测试集上反复验证。
  4. 处理速度慢

    • 原因:EAST和CRNN都是神经网络,在CPU上运行较慢。图片太大时,EAST的预处理缩放和密集预测计算量很大。
    • 解决
      • 启用GPU:确保你的TensorFlow是GPU版本,并且CUDA/cuDNN已正确安装。
      • 图片缩放:设置合理的max_side_len。对于网络图片或扫描文档,1200-1600像素通常足够。
      • 批量推理:如果有多张图片需要处理,尽量将图片组织成批次(batch)送入模型,这能充分利用GPU的并行计算能力。对于CRNN,需要将同一批次的文本行图像填充到相同宽度。

7. 进阶思路:模型训练与自定义

跑通预训练模型只是第一步。要让这套系统在你的特定场景(如识别单据、车牌、古籍等)下表现优异,通常需要进行微调或重新训练。

7.1 数据准备与标注

  • 检测数据:需要标注图片中文本区域的四边形顶点坐标。标注工具可以使用labelImg(支持旋转矩形)、PPOCRLabelRoboflow。标注格式需要转换成模型代码要求的格式,通常是每行一个文本框,坐标用逗号分隔,如x1,y1,x2,y2,x3,y3,x4,y4,transcript( transcript 对于检测任务可以置为 ‘###’ 或忽略)。
  • 识别数据:需要大量的单行文本图像和对应的文本标签。图像可以从检测数据中裁剪得到,标签就是框内的文字。注意,识别数据需要处理各种字体、大小、背景和扭曲情况。

7.2 训练EAST/AdvancedEAST

  1. 数据生成:大部分开源代码会提供数据生成脚本,将标注文件转换为训练所需的格式(如TFRecord)。
  2. 损失函数:EAST的损失由两部分加权组成:文本分类损失(通常使用平衡交叉熵)和几何形状回归损失(对于RBOX是IoU损失或平滑L1损失,对于QUAD是顶点坐标的平滑L1损失)。
  3. 训练技巧
    • 数据增强:至关重要。包括随机旋转、缩放、裁剪、颜色抖动、模糊、弹性变换等,以增加模型鲁棒性。
    • 学习率策略:使用余弦退火或带热重启的余弦退火(CosineAnnealingWarmRestarts)通常效果不错。
    • 主干网络:可以尝试更强大的主干,如ResNet50、EfficientNet,但要注意计算量和预训练权重的适配。

7.3 训练CRNN

  1. 字符表定义:收集所有训练集中出现的字符,生成char_list
  2. CTC解码器:训练时使用CTC损失,无需对齐。Keras/TensorFlow中可以使用tf.keras.backend.ctc_batch_cost自定义损失函数。
  3. 训练技巧
    • 输入图像归一化:保持一致。
    • 序列长度:CRNN可以处理可变宽度,但训练时一个批次内的图像需要填充到相同宽度。可以使用tf.data.Datasetpadded_batch方法。
    • 使用注意力机制:在CRNN的RNN部分之后加入注意力机制(如Bahdanau Attention),可以帮助模型更好地聚焦于相关字符区域,对复杂背景或艺术字体有奇效。但这会改变模型结构,需要调整解码部分。

我个人在训练中的体会是,数据质量远比模型结构重要。一个在干净数据集上训练的简单模型,往往比在嘈杂数据上训练的复杂模型在实际应用中更可靠。对于工业场景,花大量时间清洗和扩增训练数据,是最有价值的投资。

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/8/28 10:11:53

用Python和Skyfield实现地址查询日食可见度

最近在 Hacker News 上看到一个很有意思的 Show HN 作品&#xff1a;输入一个地址&#xff0c;页面就会告诉你 8 月 12 日的日食在你家屋顶上看起来是什么效果。这类工具平常看起来只是“地图 天文数据”的简单拼接&#xff0c;但真正实现时&#xff0c;你会发现地址解析、天文…

作者头像 李华
网站建设 2026/8/28 10:09:14

基于Nordic BLE SoC的资产追踪系统设计与低功耗调优实战

1. 项目缘起 做了这么多年物联网硬件&#xff0c;我越来越觉得&#xff0c;资产追踪这个赛道像是被低估的宝藏。不少团队一上来就盯着GPS、4G Cat.1或者LoRa&#xff0c;觉得覆盖远、信号强才是王道。但真到了实际项目里——尤其是室内仓储、园区设备盘点、工具借还管理、医疗设…

作者头像 李华
网站建设 2026/8/28 10:09:12

高斯整数与费马平方和定理:解决圆上整点问题的数论模板

1. 项目概述&#xff1a;从一道题到一类问题的解法 最近在洛谷上刷题&#xff0c;又碰到了那道经典的“圆上的整点”&#xff08;P2508&#xff09;。题目本身描述很简单&#xff1a;给定一个正整数 $n$&#xff0c;求以原点为圆心、以 $\sqrt{n}$ 为半径的圆上&#xff0c;有多…

作者头像 李华
网站建设 2026/8/28 10:04:28

172、夜景模式的多帧合成芯片级优化——基于海思ISP的RAW域对齐与降噪的NPU卸载与功耗控制

172、夜景模式的多帧合成芯片级优化——基于海思ISP的RAW域对齐与降噪的NPU卸载与功耗控制 去年在珠海做某旗舰机型的夜景模组调试,客户反馈一个非常诡异的现象:夜景模式连续拍十张,前五张画质完美,第六张开始画面出现微妙的“糊”,到第八张直接出现重影。我们一开始怀疑…

作者头像 李华
网站建设 2026/8/28 10:04:01

ComfyUI工作流合集:AI绘画效率提升与节点式工作流实战指南

简介&#xff1a;在AI绘画领域&#xff0c;Stable Diffusion等扩散模型通过将文本描述转化为高质量图像&#xff0c;已成为内容创作的重要工具。其核心原理是基于潜在空间的迭代去噪过程&#xff0c;通过提示词引导生成方向。节点式工作流作为实现这一过程的可视化编程界面&…

作者头像 李华