【pipeline实践】基于mmocr和文本合成器快速训练一个自己的文本识别器

mmocr的安装

  mmocr的安装实测下来没有什么坑,按照官方文档的指导可以快速、顺利的安装完毕:

conda create -n open-mmlab python=3.8 pytorch=1.10 cudatoolkit=11.3 torchvision -c pytorch -y
conda activate open-mmlab
pip3 install openmim
git clone https://github.com/open-mmlab/mmocr.git
cd mmocr
mim install -e .

私有数据准备

   mm官方提供了文本识别领域最为经典也是在产业界落地最为广泛的CRNN算法的训练model供用户下载、试用。该模型是基于码表lower_english_digits.txt,训练数据使用的是合成数据集Syn90k。
  但该公开模型可能仅供试玩,如果是真实应用场景,大概率要准备自己的文本数据,训练自己的crnn模型。在这种情况下,数据集的形式应该准备为如下形式:

...
{"filename": "20250428_204448_09.png", "text": "09"}
{"filename": "20250428_204448_1.png", "text": "1"}
{"filename": "20250428_204448_8059980.png", "text": "8059980"}
...

存储在textrecog_train.jsonl一个jsonl文件中。图片则放在与该jsonl相同的目录下。
  接下来是数据配置文件和模型配置文件,关于数据配置文件,可以增添一个对应的数据py配置文件,例如configs/textrecog/base/datasets/simplesynth_wireone.py, 这里要重点注意type为’RecogTextDataset’而非’OCRDataset’

wireone_textrecog_data_root = 'data/wireone'

wireone_textrecog_train = dict(
    type='RecogTextDataset',
    data_root=wireone_textrecog_data_root,
    ann_file='textrecog_train.jsonl',
    pipeline=None)

  然后增添一个模型训练配置文件configs/textrecog/crnn/crnn_simplesequence.py

# training schedule for 1x
_base_ = [
    '../_base_/datasets/simplesynth_wireone.py',
    '../_base_/default_runtime.py',
    '../_base_/schedules/schedule_adadelta_5e.py',
    '_base_crnn_mini-vgg.py',
]
# dataset settings
train_list = [_base_.wireone_textrecog_train]

test_list = [_base_.wireone_textrecog_train]

default_hooks = dict(logger=dict(type='LoggerHook', interval=50), )
train_dataloader = dict(
    batch_size=64,
    num_workers=24,
    persistent_workers=True,
    sampler=dict(type='DefaultSampler', shuffle=True),
    dataset=dict(
        type='ConcatDataset',
        datasets=train_list,
        pipeline=_base_.train_pipeline))
test_dataloader = dict(
    batch_size=1,
    num_workers=4,
    persistent_workers=True,
    drop_last=False,
    sampler=dict(type='DefaultSampler', shuffle=False),
    dataset=dict(
        type='ConcatDataset',
        datasets=test_list,
        pipeline=_base_.test_pipeline))
val_dataloader = test_dataloader

val_evaluator = dict(
    dataset_prefixes=['wireone'])
test_evaluator = val_evaluator

auto_scale_lr = dict(base_batch_size=64 * 4)

模型的训练与推理

  在完成上述数据和配置文件的准备,模型的训练执行如下指令:

python tools/train.py configs/textrecog/crnn/crnn_simplesequence.py

  推理指令

python tools/infer.py demo/20250428_211548_197705.png --rec configs/textrecog/crnn/crnn_simplesequence.py --rec-weights work_dirs/crnn_simplesequence/epoch_50.pth --print-result
{'predictions': [{'rec_texts': ['197705'], 'rec_scores': [1.0]}]}

onnx模型的导出与验证

  基于模型训练时的配置文件和训练得到的模型文件,利用如下的脚本可以得到导出的模型。

import torch
from mmocr.apis.inferencers import MMOCRInferencer

rec = "../../whatever/mmocr/configs/textrecog/crnn/crnn_simplesequence.py"
rec_weights = "../../whatever/mmocr/work_dirs/crnn_simplesequence/epoch_50.pth"
ocr_inferencer = MMOCRInferencer(rec=rec, rec_weights=rec_weights)
torch.onnx.export(ocr_inferencer.textrec_inferencer.model, (torch.randn(1, 1, 32, 48, device=next(ocr_inferencer.textrec_inferencer.model.parameters()).device),), "crnn_simplesequence.onnx", input_names=["input"], dynamic_axes={'input' : {0 : 'batch_size', 3: "width"}})

onnx模型在c++推理库

SequenceInference

反思mmocr

  感觉mm系列的仓库发展到现在,对于engineer来讲,像是传统图像处理领域的opencv。engineer无需了解内部代码、也不希望engineer修改内部代码。使用者仅仅需要了解接口参数也即配置参数的含义,即可快速的获得一个可用的模型。对于researcher来讲则是一个快速复现对比算法指标的toolbox,另外一方面对于算法Pipeine模块化的划分,也有助于researcher做基于当前已有pipeline做模块化级别的创新研究。

更多推荐