PaddleOCR 模型微调实战:基于 PP-OCRv3 检测与识别模型的垂类场景精度提升指南
【免费下载链接】PaddleOCRTurn any PDF or image document into structured data for your AI. A powerful, lightweight OCR toolkit that bridges the gap between images/PDFs and LLMs. Supports 100+ languages.项目地址: https://gitcode.com/GitHub_Trending/pa/PaddleOCR
PaddleOCR 官方提供的 PP-OCR 系列模型在通用场景下已具备出色的检测与识别能力,但当业务落地到票据、证照、仪表盘、数码管等垂类场景时,通过模型微调(Fine-tune)可以进一步显著提升精度。本文以 PaddleOCR 仓库中 模型微调文档 为骨架,系统讲解文本检测与文本识别模型微调的数据准备、模型与超参选择、预测参数调优及迭代训练方法,并结合仓库源码与配置逐项印证,帮助读者在自己的数据集上训练出精度更高、泛化更好的检测与识别模型。
1. 模型微调背景与意义
PP-OCR 系列模型在通用场景中性能优异,能够解决绝大多数情况下的检测与识别问题。但在垂类场景中,如果希望获取更优的模型效果,可以通过模型微调的方法,进一步提升 PP-OCR 系列检测与识别模型的精度。
文本检测与识别模型的微调流程整体一致:下载官方预训练模型 → 准备垂类数据 → 调整配置中的预训练路径与超参 → 启动训练 → 评估与 badcase 分析 → 迭代优化。整个流程中需要把握以下核心要点:
- PP-OCR 提供的预训练模型有较好的泛化能力,是微调的理想起点;
- 加入少量真实数据(检测任务 ≥500 张、识别任务 ≥5000 张),会大幅提升垂类场景的检测与识别效果;
- 在模型微调时,加入真实通用场景数据,可以进一步提升模型精度与泛化性能;
- 在文本检测任务中,增大图像的预测尺度,能够进一步提升较小文字区域的检测效果;
- 在模型微调时,需要适当调整超参数(学习率、batch size 最为重要),以获得更优的微调效果。
2. 文本检测模型微调
2.1 数据选择
- 数据量:建议至少准备 500 张文本检测数据集用于模型微调。少量高质量的真实数据即可大幅提升垂类场景的检测效果。
- 数据标注:采用单行文本标注格式,建议标注的检测框与实际语义内容一致。例如在火车票场景中,姓氏与名字可能离得较远,但它们在语义上属于同一个检测字段,这里也需要将整个姓名标注为 1 个检测框。
检测任务的标注格式与数据准备细节可参考 OCR 数据集文档 以及 PP-OCRv3 文本检测模型训练文档 中的数据准备章节。
2.2 模型选择
建议选择 PP-OCRv3 检测模型进行微调:
- 配置文件:configs/det/PP-OCRv3/PP-OCRv3_mobile_det.yml;
- 预训练模型:
ch_PP-OCRv3_det_distill_train.tar,解压后得到ch_PP-OCRv3_det_distill_train文件夹; - 更多 PP-OCR 系列模型可参考 PP-OCR 系列模型库。
注意:
ch_PP-OCRv3_det_distill_train.tar是使用 CML(Collaborative Mutual Learning)协同互学习蒸馏策略训练得到的产物,压缩包内同时包含 Student、Student2、Teacher 三份参数。在使用上述预训练模型时,必须使用文件夹中的student.pdparams文件作为预训练模型,即仅使用学生模型。
关于如何从蒸馏模型压缩包中提取student.pdparams(通过paddle.load加载全部参数后,以key[len("Student."):]过滤并paddle.save保存),可参考 PP-OCRv3 文本检测模型训练文档 中"提取 Student 参数"一节,其中给出了完整的 Python 脚本。
2.3 训练超参选择
在模型微调时,最重要的超参是预训练模型路径pretrained_model、学习率learning_rate与batch_size。检测微调的部分关键配置如下:
Global: pretrained_model: ./ch_PP-OCRv3_det_distill_train/student.pdparams # 预训练模型路径 Optimizer: lr: name: Cosine learning_rate: 0.001 # 学习率 warmup_epoch: 2 regularizer: name: 'L2' factor: 0 Train: loader: shuffle: True drop_last: False batch_size_per_card: 8 # 单卡batch size num_workers: 4使用上述配置时,首先需要将pretrained_model字段指定为student.pdparams文件路径(配置为./ch_PP-OCRv3_det_distill_train/student或指向解压后的完整路径均可,PaddleOCR 会自动补全.pdparams后缀)。
需要特别说明的是:PaddleOCR 提供的配置文件是在 8 卡训练(相当于总的 batch size 是8*8=64)、且没有加载预训练模型情况下给出的默认配置,因此在实际场景中,学习率需要与总的 batch size 进行对应线性调整。例如:
- 如果您的场景中是单卡训练,单卡 batch_size=8,则总的 batch_size=8,建议将学习率调整为
1e-4左右; - 如果您的场景中是单卡训练,由于显存限制,只能设置单卡 batch_size=4,则总的 batch_size=4,建议将学习率调整为
5e-5左右。
参考仓库中 configs/det/PP-OCRv3/PP-OCRv3_mobile_det.yml 的原始配置,其学习率调度器为Cosine,初始学习率0.001,warmup_epoch: 2,权重衰减使用 L2 正则、factor 为5.0e-05。学习率调度器的具体实现可参见 ppocr/optimizer/learning_rate.py 中的Cosine、LinearWarmupCosine、Piecewise等类。
启动检测微调训练的命令如下(-o用于覆盖配置项,无需手工修改 yml 文件):
# 单卡训练 python3 tools/train.py -c configs/det/PP-OCRv3/PP-OCRv3_mobile_det.yml \ -o Global.pretrained_model=./student \ Global.save_model_dir=./output/ # 多卡分布式训练 python3 -m paddle.distributed.launch --gpus '0,1,2,3' tools/train.py \ -c configs/det/PP-OCRv3/PP-OCRv3_mobile_det.yml \ -o Global.pretrained_model=./student \ Global.save_model_dir=./output/训练入口 tools/train.py 会依次完成:读取并解析配置文件 → 构建数据加载器(build_dataloader)→ 构建后处理(build_post_process)→ 构建模型(build_model)→ 构建损失(build_loss)→ 构建优化器与学习率调度器(build_optimizer)→ 构建评估指标(build_metric)→ 加载预训练模型(load_model)→ 调用program.train启动训练。其中pretrained_model正是通过load_model加载的,微调时务必确认该路径正确指向解压后的权重文件。
2.4 预测超参选择
对训练好的模型导出并进行推理时,可以通过进一步调整预测的图像尺度,来提升小面积文本的检测效果。以 DBNet 推理为例,下面这些超参数可以通过适当调整来提升效果:
| 参数名称 | 类型 | 默认值 | 含义 |
|---|---|---|---|
| det_db_thresh | float | 0.3 | DB输出的概率图中,得分大于该阈值的像素点才会被认为是文字像素点 |
| det_db_box_thresh | float | 0.6 | 检测结果边框内,所有像素点的平均得分大于该阈值时,该结果会被认为是文字区域 |
| det_db_unclip_ratio | float | 1.5 | Vatti clipping算法的扩张系数,使用该方法对文字区域进行扩张 |
| max_batch_size | int | 10 | 预测的 batch size |
| use_dilation | bool | False | 是否对分割结果进行膨胀以获取更优检测效果 |
| det_db_score_mode | str | "fast" | DB的检测结果得分计算方法,支持fast和slow,fast是根据 polygon 的外接矩形边框内的所有像素计算平均得分,slow是根据原始 polygon 内的所有像素计算平均得分,计算速度相对较慢一些,但是更加准确一些 |
这些参数在仓库源码 ppocr/postprocess/db_postprocess.py 的DBPostProcess类中均有对应实现:thresh对应二值化阈值(segmentation = pred > self.thresh),box_thresh用于过滤得分过低的候选框(if self.box_thresh > score: continue),unclip_ratio决定文字区域扩张幅度(self.unclip(points, self.unclip_ratio),扩张距离为poly.area * unclip_ratio / poly.length),score_mode则对应box_score_fast(外接矩形内像素平均分)与box_score_slow(原始多边形内像素平均分)两种得分计算方式。
更多关于推理方法的介绍可以参考 Paddle Inference 推理教程。
3. 文本识别模型微调
3.1 数据选择
- 数据量:不更换字典的情况下,建议至少准备 5000 张的文本识别数据集用于模型微调;如果更换了字典(不建议),需要的数量更多。
- 数据分布:建议分布与实测场景尽量一致。如果实测场景包含大量短文本,则训练数据中建议也包含较多短文本;如果实测场景对于空格识别效果要求较高,则训练数据中建议也包含较多带空格的文本内容。
- 数据合成:针对部分字符识别有误的情况,建议获取一批特定字符数据,加入到原数据中使用小学习率微调。其中原始数据与新增数据比例可尝试10:1 ~ 5:1,避免单一场景数据过多导致模型过拟合,同时尽量平衡语料词频,确保常用字的出现频率不会过低。特定字符可以使用TextRenderer等合成工具生成(仓库中多语言数据集的合成数据即使用了开源合成工具 text_renderer,可参考 数据合成文档),合成数据语料尽量来自真实使用场景,在贴近真实场景的基础上保持字体、背景的丰富性,有助于提升模型效果。
- 通用中英文数据:在训练的时候,可以在训练集中添加通用真实数据(如在不更换字典的微调场景中,建议添加 LSVT、RCTW、MTWI 等真实数据),进一步提升模型的泛化性能。
识别数据集的目录组织、标注文件格式(图片路径\t标注内容逐行写入 txt)、字典文件格式与内置字典列表(如 ppocr/utils/ppocr_keys_v1.txt 为包含 6623 个字符的中文字典),可参考 文字识别文档 的"数据准备"章节。
3.2 模型选择
建议选择 PP-OCRv3 识别模型进行微调:
- 配置文件:configs/rec/PP-OCRv3/PP-OCRv3_mobile_rec_distillation.yml;
- 预训练模型:
ch_PP-OCRv3_rec_train.tar,解压后使用其中的ch_PP-OCRv3_rec_train/best_accuracy.pdparams; - 更多 PP-OCR 系列模型可参考 PP-OCR 系列模型库。
关键点:去除 GTC 策略。PP-OCRv3 模型使用了 GTC(Guided Training of CTC)策略,其 SAR 分支参数量大,当训练数据为简单场景时模型容易过拟合,导致微调效果不佳,因此建议在微调时去除 GTC 策略。从 configs/rec/PP-OCRv3/PP-OCRv3_mobile_rec_distillation.yml 可以看到,原模型结构为DistillationModel,其中每个分支的 Head 均为MultiHead(同时包含 CTCHead 与 SARHead),并配合DistillationDMLLoss、DistillationSARLoss等蒸馏损失;微调时建议将模型结构修改为单一 SVTR + CTCHead的简化结构,同时将标签编码方式从MultiLabelEncode改回CTCLabelEncode,并去除RecConAug增广,模型结构部分配置修改如下:
Architecture: model_type: rec algorithm: SVTR Transform: Backbone: name: MobileNetV1Enhance scale: 0.5 last_conv_stride: [1, 2] last_pool_type: avg Neck: name: SequenceEncoder encoder_type: svtr dims: 64 depth: 2 hidden_dims: 120 use_guide: False Head: name: CTCHead fc_decay: 0.00001 Loss: name: CTCLoss Train: dataset: ...... transforms: # 去除 RecConAug 增广 # - RecConAug: # prob: 0.5 # ext_data_num: 2 # image_shape: [48, 320, 3] # max_text_length: *max_text_length - RecAug: # 修改 Encode 方式 - CTCLabelEncode: - KeepKeys: keep_keys: - image - label - length ... Eval: dataset: ... transforms: ... - CTCLabelEncode: - KeepKeys: keep_keys: - image - label - length ...说明:由于原配置文件是蒸馏配置(
Architecture.algorithm: Distillation,Loss 为CombinedLoss,包含 DistillationDMLLoss / DistillationDistanceLoss / DistillationCTCLoss / DistillationSARLoss),微调时若直接沿用该配置则必须保留完整的 Teacher/Student 双分支结构与best_accuracy.pdparams中的全部参数;而本节给出的简化结构配置去除了 SAR 分支与蒸馏损失,可有效规避小数据量场景下的过拟合,是官方推荐的垂类微调做法。
3.3 训练超参选择
与文本检测任务微调相同,在识别模型微调时,最重要的超参同样是预训练模型路径pretrained_model、学习率learning_rate与batch_size。识别模型微调的部分默认配置如下:
Global: pretrained_model: # 预训练模型路径 Optimizer: lr: name: Piecewise decay_epochs : [700, 800] values : [0.001, 0.0001] # 学习率 warmup_epoch: 5 regularizer: name: 'L2' factor: 0 Train: dataset: name: SimpleDataSet data_dir: ./train_data/ label_file_list: - ./train_data/train_list.txt ratio_list: [1.0] # 采样比例,默认值是[1.0] loader: shuffle: True drop_last: False batch_size_per_card: 128 # 单卡batch size num_workers: 8使用上述配置时,首先需要将pretrained_model字段指定为 3.2 章节中解压得到的ch_PP-OCRv3_rec_train/best_accuracy.pdparams文件路径。
同样地,PaddleOCR 提供的配置文件是在 8 卡训练(相当于总的 batch size 是8*128=1024)、且没有加载预训练模型情况下给出的默认配置,因此您的场景中学习率与总的 batch size 需要对应线性调整,例如:
- 如果您的场景中是单卡训练,单卡 batch_size=128,则总的 batch_size=128,在加载预训练模型的情况下,建议将学习率调整为
[1e-4, 2e-5]左右(Piecewise 学习率策略需设置 2 个值,对应values字段); - 如果您的场景中是单卡训练,因为显存限制,只能设置单卡 batch_size=64,则总的 batch_size=64,在加载预训练模型的情况下,建议将学习率调整为
[5e-5, 1e-5]左右。
Piecewise学习率调度器会在decay_epochs指定的 epoch 处将学习率切换为values中对应的值,其实现同样位于 ppocr/optimizer/learning_rate.py。原仓库配置中识别模型的默认值为decay_epochs: [700]、values: [0.0005, 0.00005],微调时可根据训练总 epoch 数重新设定衰减节点。
通用数据与垂类数据的配比:如果有通用真实场景数据加进来,建议每个 epoch 中,垂类场景数据与真实场景的数据量保持在1:1左右。例如:您自己的垂类场景识别数据量为 1W,数据标签文件为vertical.txt;收集到的通用场景识别数据量为 10W,数据标签文件为general.txt。那么可以设置label_file_list和ratio_list参数如下所示:
Train: dataset: name: SimpleDataSet data_dir: ./train_data/ label_file_list: - vertical.txt - general.txt ratio_list: [1.0, 0.1]这样配置后,每个 epoch 中vertical.txt会进行全采样(采样比例为 1.0),包含 1W 条数据;general.txt会按照 0.1 的采样比例进行采样,包含10W*0.1=1W条数据,最终二者的比例为1:1。
该采样机制的底层实现位于 ppocr/data/simple_dataset.py(lmdb_dataset.py、pubtab_dataset.py、pgnet_dataset.py中也有相同逻辑):ratio_list的长度必须与label_file_list一致,训练模式下每个 epoch 会按count = round(file_size * ratio_list[i])对该文件随机采样count条数据,当任一文件的ratio_list小于 1.0 时,数据加载器会在每个 epoch 间进行need_reset重采样,从而实现多数据源的按比例混合。
3.4 训练调优
训练过程并非一蹴而就的,完成一个阶段的训练评估后,建议收集分析当前模型在真实场景中的 badcase,有针对性地调整训练数据比例,或者进一步新增合成数据。通过多次迭代训练,不断优化模型效果。
关于训练日志字段解读(epoch / iter / lr / loss / acc / norm_edit_dis / ips 等)、评估命令(tools/eval.py+-o Global.checkpoints=...)、断点训练(Global.checkpoints优先级高于Global.pretrained_model)、模型导出(tools/export_model.py+Global.save_inference_dir)以及混合精度训练(Global.use_amp=True)等完整操作流程,可参考 文字识别文档 与 PP-OCRv3 文本检测模型训练文档 的对应章节。
另外有一个常见现象需要提前知晓:如果在训练时修改了自定义字典,由于无法加载最后一层 FC(全连接输出层)的参数,在迭代初期 acc=0 是正常的情况,不必担心,加载预训练模型依然可以加快模型收敛。从 tools/train.py 的源码可以看到,训练时会根据后处理解析出的字符数量动态设置模型 Head 的out_channels(即char_num),修改字典后输出维度随之变化,最后一层参数自然无法与预训练权重对齐,因此初期的 acc=0 属于预期行为,随迭代进行会逐步恢复正常。
4. 微调完整流程速查
综合上述内容,一次完整的 PaddleOCR 模型微调可按以下步骤执行:
- 准备数据:检测任务准备 ≥500 张已标注垂类图片;识别任务准备 ≥5000 张(不更换字典)图片,标注文件为
图片路径\t标签逐行格式,存放于train_data/目录; - 下载预训练模型:检测模型下载
ch_PP-OCRv3_det_distill_train.tar并解压,提取student.pdparams;识别模型下载ch_PP-OCRv3_rec_train.tar并解压; - 修改配置:将
Global.pretrained_model指向预训练权重;按实际总 batch size 线性缩放学习率(检测默认总 batch size=64、识别默认总 batch size=1024 为基准);识别任务建议按 3.2 节去除 GTC 策略并改回CTCLabelEncode;如需混入通用数据,配置label_file_list与ratio_list保持 1:1 配比; - 启动训练:
python3 tools/train.py -c <配置文件> -o Global.pretrained_model=... Global.save_model_dir=./output/,多卡时使用python3 -m paddle.distributed.launch --gpus '0,1,2,3' ...; - 评估与预测调优:用
tools/eval.py评估;推理时通过det_db_thresh、det_db_box_thresh、det_db_unclip_ratio、det_db_score_mode、use_dilation等 DBNet 后处理参数调节小字检测效果; - 迭代优化:收集真实场景 badcase,调整数据配比或补充合成数据,反复迭代直至精度收敛。
通过以上方法,即可在 PP-OCRv3 预训练模型的基础上,用少量真实数据快速获得面向自身垂类场景的高精度检测与识别模型。
【免费下载链接】PaddleOCRTurn any PDF or image document into structured data for your AI. A powerful, lightweight OCR toolkit that bridges the gap between images/PDFs and LLMs. Supports 100+ languages.项目地址: https://gitcode.com/GitHub_Trending/pa/PaddleOCR
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考