news 2026/9/11 2:44:17

PaddleOCR 文本检测模型训练、评估与推理全流程实战:以 icdar2015 为例

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PaddleOCR 文本检测模型训练、评估与推理全流程实战:以 icdar2015 为例

PaddleOCR 文本检测模型训练、评估与推理全流程实战:以 icdar2015 为例

【免费下载链接】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 仓库中的 检测模型训练文档 为核心骨架,结合 tools/train.py、tools/eval.py、tools/export_model.py 等源码与 det_mv3_db.yml 配置逐项对照验证,系统讲解从数据准备、权重下载、训练启动(单卡/多卡/多机/混合精度)、评估指标、单图/批量测试到推理模型导出的完整链路。

1. 数据与预训练权重准备

1.1 数据准备

PaddleOCR 的检测模型训练以icdar2015数据集作为官方示例。数据集的下载、标注格式与目录组织方式请参考 OCR 数据集文档。训练集与验证集的标注文件分别对应配置文件中的Train.dataset.label_file_listEval.dataset.label_file_list,例如:

Train: dataset: name: SimpleDataSet data_dir: ./train_data/icdar2015/text_localization/ label_file_list: - ./train_data/icdar2015/text_localization/train_icdar2015_label.txt Eval: dataset: name: SimpleDataSet data_dir: ./train_data/icdar2015/text_localization/ label_file_list: - ./train_data/icdar2015/text_localization/test_icdar2015_label.txt

在 det_mv3_db.yml 中可以看到Train.dataset使用SimpleDataSet,通过data_dir + label_file_list定位图片与标签;Eval.dataset则指向test_icdar2015_label.txt,且Eval.loader.batch_size_per_card必须为1(评估过程逐图进行)。

1.2 下载预训练骨干权重

PaddleOCR 的检测模型目前支持 3 种骨干网络:MobileNetV3、ResNet18_vd、ResNet50_vd。预训练权重统一放到./pretrain_models/目录下:

cd PaddleOCR/ # 下载 MobileNetV3 预训练模型 wget -P ./pretrain_models/ https://paddleocr.bj.bcebos.com/pretrained/MobileNetV3_large_x0_5_pretrained.pdparams # 或下载 ResNet18_vd 预训练模型 wget -P ./pretrain_models/ https://paddleocr.bj.bcebos.com/pretrained/ResNet18_vd_pretrained.pdparams # 或下载 ResNet50_vd 预训练模型 wget -P ./pretrain_models/ https://paddleocr.bj.bcebos.com/pretrained/ResNet50_vd_ssld_pretrained.pdparams

说明:预训练权重下载后,训练时通过Global.pretrained_model指定路径(不带.pdparams后缀,如./pretrain_models/MobileNetV3_large_x0_5_pretrained)。这些骨干权重来自 PaddleClas 的分类预训练模型,仅用于初始化特征提取层。

2. 训练

2.1 启动训练

使用tools/train.py启动训练,-c指定配置文件,-o用于覆盖配置项:

python3 tools/train.py -c configs/det/det_mv3_db.yml \ -o Global.pretrained_model=./pretrain_models/MobileNetV3_large_x0_5_pretrained

若安装的是 CPU 版本 PaddlePaddle,请将配置中的use_gpu设为false

-o支持任意层级的键值覆盖,无需修改 yml 文件。例如将学习率调整为 0.0001:

# 单 GPU 训练 python3 tools/train.py -c configs/det/det_mv3_db.yml -o \ Global.pretrained_model=./pretrain_models/MobileNetV3_large_x0_5_pretrained \ Optimizer.base_lr=0.0001 # 多 GPU 训练:通过 --gpus 指定使用的 GPU ID python3 -m paddle.distributed.launch --gpus '0,1,2,3' tools/train.py -c configs/det/det_mv3_db.yml \ -o Global.pretrained_model=./pretrain_models/MobileNetV3_large_x0_5_pretrained # 多机多卡训练:通过 --ips 指定节点 IP,--gpus 指定 GPU ID python3 -m paddle.distributed.launch --ips="xx.xx.xx.xx,xx.xx.xx.xx" --gpus '0,1,2,3' \ tools/train.py -c configs/det/det_mv3_db.yml \ -o Global.pretrained_model=./pretrain_models/MobileNetV3_large_x0_5_pretrained

多机训练注意事项--ips必须替换为各机器的实际地址,且机器之间需能互相 ping 通;需要在多台机器上分别启动训练命令,查看本机 IP 可用ifconfig

想要进一步加速训练,可开启自动混合精度训练。单卡训练命令如下:

python3 tools/train.py -c configs/det/det_mv3_db.yml \ -o Global.pretrained_model=./pretrain_models/MobileNetV3_large_x0_5_pretrained \ Global.use_amp=True Global.scale_loss=1024.0 Global.use_dynamic_loss_scaling=True

从源码看训练主流程:在 tools/train.py 中,main()依次完成分布式环境初始化(dist.init_parallel_env)、构建训练/验证 DataLoader(build_dataloader)、构建后处理(build_post_process)、构建模型(build_model)、构建损失(build_loss)、构建优化器(build_optimizer)与评估指标(build_metric),最后调用program.train()进入训练循环。use_amp开启时,会在 tools/train.py 中构造paddle.amp.GradScaler,并按amp_level(默认O2)对模型与优化器进行paddle.amp.decorate封装,同时设置master_weight=True保证主权重精度。

2.2 加载已训练模型继续训练

若希望加载训练中间产物(checkpoints)断点续训,指定Global.checkpoints即可:

python3 tools/train.py -c configs/det/det_mv3_db.yml -o Global.checkpoints=./your/trained/model

注意Global.checkpoints的优先级高于Global.pretrained_model;当两者同时指定时,优先加载Global.checkpoints指向的模型;若该路径错误,则回退加载Global.pretrained_model指向的模型。这一加载逻辑由 ppocr/utils/save_load.py 中的load_model实现。

2.3 使用新骨干网络训练

PaddleOCR 将检测网络划分为四个串联模块,数据依次经过transforms -> backbones -> necks -> heads,相关代码位于 ppocr/modeling 目录:

├── architectures # 网络构建代码 ├── transforms # 图像变换模块 ├── backbones # 特征提取模块 ├── necks # 特征增强模块 └── heads # 输出模块

如果目标骨干在 PaddleOCR 中已有实现,直接修改配置文件Backbone部分即可;若需引入全新的 Backbone,步骤如下:

  1. 在 ppocr/modeling/backbones 目录下新建文件,例如my_backbone.py
  2. 在其中编写继承paddle.nn.Layer的网络类:
import paddle import paddle.nn as nn import paddle.nn.functional as F class MyBackbone(nn.Layer): def __init__(self, *args, **kwargs): super(MyBackbone, self).__init__() # 你的初始化代码 self.conv = nn.xxxx def forward(self, inputs): # 你的网络前向逻辑 y = self.conv(inputs) return y
  1. 在 ppocr/modeling/backbones/init.py 中导入新模块。

四个模块添加完成后,只需在配置文件中声明即可使用:

Backbone: name: MyBackbone args1: args1

说明:替换 Backbone 及其他模块的完整规范见 新增算法文档。从配置看,Architecture采用模块化注册机制(如DBFPNNeck、DBHeadHead、DBLossLoss、DBPostProcess后处理),各模块通过name字段在对应目录的__init__.py中完成注册与实例化。

2.4 混合精度训练

希望进一步加速训练时,可使用自动混合精度训练。以单机单卡为例:

python3 tools/train.py -c configs/det/det_mv3_db.yml \ -o Global.pretrained_model=./pretrain_models/MobileNetV3_large_x0_5_pretrained \ Global.use_amp=True Global.scale_loss=1024.0 Global.use_dynamic_loss_scaling=True

其中Global.scale_loss为梯度缩放初始值(init_loss_scaling),Global.use_dynamic_loss_scaling决定是否启用动态损失缩放,二者直接传入paddle.amp.GradScaler(见 tools/train.py)。

2.5 分布式训练

多机多卡训练时,--ips指定机器 IP,--gpus指定 GPU ID:

python3 -m paddle.distributed.launch --ips="xx.xx.xx.xx,xx.xx.xx.xx" --gpus '0,1,2,3' \ tools/train.py -c configs/det/det_mv3_db.yml \ -o Global.pretrained_model=./pretrain_models/MobileNetV3_large_x0_5_pretrained

注意事项:

  1. --ips需替换为各机器实际地址,且机器之间可互相 ping 通;
  2. 需在多个机器上分别启动训练;
  3. 查看本机 IP 使用ifconfig
  4. 更多分布式训练加速比细节见 分布式训练教程。

从源码结构可以推断:配置文件中Global.distributed为真时,训练代码会先执行dist.init_parallel_env(),并将模型包装为paddle.DataParallel(见 tools/train.py 与 tools/train.py)。

2.6 知识蒸馏训练

PaddleOCR 的文本检测训练支持知识蒸馏(Knowledge Distillation),通常用于训练轻量级学生模型,细节参考 知识蒸馏文档。蒸馏配置在Architecture中声明多个子模型(teacher/student),后处理则使用DistillationDBPostProcess(见 ppocr/postprocess/db_postprocess.py),其内部为每个子模型(默认model_name=["student"])分别执行 DB 后处理。

2.7 其他平台训练(Windows / macOS / Linux DCU)

  • Windows GPU/CPU:Windows 平台仅支持单 GPU 训练与推理,用set CUDA_VISIBLE_DEVICES=0指定 GPU;DataLoader 仅支持单进程模式,需将num_workers设为 0。
  • macOS:不支持 GPU 模式,需在配置文件中将use_gpu设为False,其余训练/评估/预测命令与 Linux GPU 完全一致。
  • Linux DCU:在 DCU 设备上运行需设置环境变量export HIP_VISIBLE_DEVICES=0,1,2,3,其余命令与 Linux GPU 一致。

2.8 微调

实际业务中,推荐加载官方预训练模型并在自有数据集上微调。检测模型的微调方法详见 模型微调教程。微调核心要点:

  • 数据集至少准备500 张检测标注图,标注框需与语义内容一致(例如火车票场景中"姓+名"虽相距较远但语义同一字段,应标注为一个检测框);
  • 推荐使用 PP-OCRv3 检测模型作为预训练权重(配置文件 PP-OCRv3_mobile_det.yml,权重包解压后使用其中的student.pdparams,即仅使用学生模型);
  • 微调时最重要的三个超参数是pretrained_modellearning_ratebatch_size。PaddleOCR 官方配置面向 8 卡训练(总 batch size = 8×8=64),你的场景需按总 batch size 线性缩放学习率:单卡 batch_size=8 时建议学习率约1e-4;单卡受显存限制 batch_size=4 时建议约5e-5
  • 推理阶段可调整预测图像尺度与 DB 后处理参数来提升小文本检测效果,常用推理超参数如下表:
超参数类型默认值含义
det_db_threshfloat0.3DB 输出的概率图中,得分大于该阈值的像素被视为文本像素
det_db_box_threshfloat0.6检测结果框内所有像素平均得分大于该阈值时,才判定为文本区域
det_db_unclip_ratiofloat1.5Vatti clipping 扩张系数,用于扩张文本区域
max_batch_sizeint10batch 大小
use_dilationboolFalse是否对分割结果做膨胀以得到更优检测结果
det_db_score_modestr"fast"DB 检测结果的得分计算方式,支持fast(按多边形外接矩形内所有像素计算平均分)与slow(按原始多边形内所有像素计算平均分,速度较慢但更准确)

3. 评估

PaddleOCR 使用Precision(精确率)、Recall(召回率)、Hmean(F1 分数)三项指标评估文本检测性能。在 ppocr/metrics/det_metric.py 中,DetMetric通过DetectionIoUEvaluator逐图比对预测多边形与 GT 多边形,最终由get_metric()汇总输出precisionrecallhmean三项指标,配置项Metric.main_indicator: hmean表明以 Hmean 作为早停/模型筛选的主指标。

运行以下命令计算评估指标,结果保存在配置文件中Global.save_res_path指定的文件里:

python3 tools/eval.py -c configs/det/det_mv3_db.yml \ -o Global.checkpoints="{path/to/weights}/best_accuracy" \ PostProcess.box_thresh=0.6 PostProcess.unclip_ratio=1.5

评估要点:

  • 评估时建议设置后处理参数box_thresh=0.6unclip_ratio=1.5;若使用不同数据集/模型训练,这两个参数需要相应调整以获得更优结果;
  • 训练过程中保存的模型参数默认存放在Global.save_model_dir目录,评估时需将Global.checkpoints指向保存的参数文件(如best_accuracy);
  • 注意box_threshunclip_ratio是 DB 后处理所需参数,评估 EAST、SAST 模型时无需设置。

从源码看评估流程:tools/eval.py 构建 Eval DataLoader 与模型后,通过load_model加载Global.checkpoints指定的权重,再调用program.eval()完成推理与指标计算;program.eval()内部使用build_post_process得到的DBPostProcess对网络输出的概率图做二值化与多边形提取(见 ppocr/postprocess/db_postprocess.py),并将box_threshunclip_ratio等参数直接用于框筛选与扩张。

4. 测试

对单张图片测试检测结果:

python3 tools/infer_det.py -c configs/det/det_mv3_db.yml \ -o Global.infer_img="./doc/imgs_en/img_10.jpg" \ Global.pretrained_model="./output/det_db/best_accuracy"

测试 DB 模型时可调整后处理阈值:

python3 tools/infer_det.py -c configs/det/det_mv3_db.yml \ -o Global.infer_img="./doc/imgs_en/img_10.jpg" \ Global.pretrained_model="./output/det_db/best_accuracy" \ PostProcess.box_thresh=0.6 PostProcess.unclip_ratio=2.0

对文件夹内所有图片测试:

python3 tools/infer_det.py -c configs/det/det_mv3_db.yml \ -o Global.infer_img="./doc/imgs_en/" \ Global.pretrained_model="./output/det_db/best_accuracy"

配置文件中Global.infer_img的默认值为doc/imgs_en/img_10.jpg(见 det_mv3_db.yml),infer_det.py同时支持传入单图路径与目录路径。

5. 推理

5.1 训练模型与推理模型的区别

  • 推理模型:由paddle.jit.save保存的固化模型,模型结构与参数已全部固化为文件,便于部署与实际系统集成;
  • checkpoints 模型:训练过程中保存的参数快照,主要用于断点续训。

与 checkpoints 相比,推理模型额外保存了模型结构信息,因此部署更简单。

5.2 导出推理模型

将 DB 训练模型转换为推理模型:

python3 tools/export_model.py -c configs/det/det_mv3_db.yml \ -o Global.pretrained_model="./output/det_db/best_accuracy" \ Global.save_inference_dir="./output/det_db_inference/"

从源码看,tools/export_model.py 通过ArgsParser解析参数、load_config加载配置、merge_config合并-o覆盖项,最终调用ppocr.utils.export_model.export(config)完成模型固化。

5.3 推理模型预测

python3 tools/infer/predict_det.py --det_algorithm="DB" \ --det_model_dir="./output/det_db_inference/" \ --image_dir="./doc/imgs/" --use_gpu=True

若使用其他检测算法(如 EAST),修改det_algorithm参数即可(默认为 DB):

python3 tools/infer/predict_det.py --det_algorithm="EAST" \ --det_model_dir="./output/det_db_inference/" \ --image_dir="./doc/imgs/" --use_gpu=True

6. FAQ

Q1:训练模型与推理模型的预测结果不一致?

A:绝大多数情况是由训练模型预测时的预处理/后处理参数与推理模型预测时的参数不一致导致。以det_mv3_db.yml训练出的模型为例,排查思路如下:

  • 检查预处理是否一致:对比 训练模型预处理配置(Eval.dataset.transforms中的DetResizeForTest,其image_shape: [736, 1280])与推理模型的预测预处理函数。评估时输入图像尺寸会影响精度——为与论文一致,icdar15 训练配置将图像 resize 到[736, 1280];而推理模型预测时只有一套默认参数,出于速度考虑默认将图像最长边限制为 960 进行 resize。两者的预处理函数均位于 ppocr/data/imaug/operators.py(DetResizeForTest等算子)。
  • 检查后处理是否一致:对比 训练模型后处理配置(PostProcess中的threshbox_threshunclip_ratio等)与推理的后处理参数是否一致。

附:检测配置深度解读(det_mv3_db.yml)

原文档对应的示例配置 configs/det/det_mv3_db.yml 是理解整个训练流程的关键,各模块与源码的对应关系如下:

配置模块关键参数对应源码/说明
Globaluse_gpuepoch_num: 1200save_model_direval_batch_step: [0, 2000]pretrained_modelcheckpointsuse_amp训练全局控制;eval_batch_step表示每 2000 个 iteration 评估一次(tools/train.py 中由program.train执行)
Architecturemodel_type: detalgorithm: DBBackbone: MobileNetV3(scale=0.5, large)Neck: DBFPN(out_channels=256)Head: DBHead(k=50)四段式网络组装(transforms→backbones→necks→heads)
LossDBLossalpha: 5beta: 10ohem_ratio: 3main_loss_type: DiceLossppocr/losses/det_db_loss.py:总损失 =alpha×shrink_map损失 + beta×threshold_map损失 + binary_map的Dice损失,其中ohem_ratio控制负样本采样比例(negative_ratio
OptimizerAdam(beta1=0.9, beta2=0.999)lr.learning_rate: 0.001L2 regularizer(factor=0)-o Optimizer.base_lr=0.0001可在线调整学习率
PostProcessDBPostProcessthresh: 0.3box_thresh: 0.6max_candidates: 1000unclip_ratio: 1.5ppocr/postprocess/db_postprocess.py:thresh为概率图二值化阈值,box_thresh为框内平均分阈值,unclip_ratio为 Vatti 扩张系数
MetricDetMetricmain_indicator: hmeanppocr/metrics/det_metric.py
Train/Eval.datasetSimpleDataSetIaaAugmentEastRandomCropDataMakeBorderMap(shrink_ratio=0.4)MakeShrinkMap等变换训练侧生成 shrink_map/threshold_map 监督信号;评估侧使用DetResizeForTest(image_shape=[736, 1280])

整体流程可以概括为:数据加载与增强(含边界/收缩图生成)→ 四段式网络前向(DB 输出 probability map、threshold map、binary map 三通道)→ DBLoss 计算(二值化可微化训练)→ 反向传播优化 → DBPostProcess 后处理提取文本框 → DetMetric 计算 Precision/Recall/Hmean → 周期性评估与模型保存 → export_model 固化推理模型 → predict_det 部署推理

【免费下载链接】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),仅供参考

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

串口协议设计六要点:从联调暴毙到稳定通信

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/11 2:39:05

CesiumJS:在浏览器里渲染 3D 地球与地图的开源库

CesiumJS:在浏览器里渲染 3D 地球与地图的开源库 【免费下载链接】cesium An open-source JavaScript library for world-class 3D globes and maps :earth_americas: 项目地址: https://gitcode.com/GitHub_Trending/ce/cesium CesiumJS 是一个基于 WebGL 的…

作者头像 李华
网站建设 2026/9/11 2:38:52

LeetCode 1224 最大相等频率:用哈希表与频率分布形态实现线性判定

LeetCode 1224 的 Maximum Equal Frequency(最大相等频率)是我刷题时印象很深的一道困难题。它名字很直白:给一个正整数数组,找出最长的一个前缀,使得我们删除前缀中的一个元素后,剩下的每个不同数字出现次…

作者头像 李华