PaddleOCR 文本检测算法 DB 与 DB++ 深度解析:可微分二值化原理、训练配置与推理部署实战
【免费下载链接】PaddleOCR飞桨多语言OCR工具包(实用超轻量OCR系统,支持80+种语言识别,提供数据标注与合成工具,支持服务器、移动端、嵌入式及IoT设备端的训练与部署) Awesome multilingual OCR toolkits based on PaddlePaddle (practical ultra lightweight OCR system, support 80+ languages recognition, provide data annotation and synthesis tools, support training and deployment among server, mobile, embedded and IoT devices)项目地址: https://gitcode.com/paddlepaddle/PaddleOCR
PaddleOCR 内置了经典的实时场景文本检测算法 DB(Differentiable Binarization)及其升级版 DB++(带自适应尺度融合 ASF),只需更换配置文件即可完成从训练、评估到推理部署的全流程。本文以 PaddleOCR 仓库中的 DB 与 DB++ 算法文档 为核心,结合仓库中的配置文件与源码实现,系统讲解两种算法的原理、复现指标、配置项含义、训练方法与多种推理部署方式,帮助读者快速在自有数据集上复现与落地该检测模型。
1. 算法简介:从 DB 到 DB++
DB(Differentiable Binarization)由 Liao Minghui 等人提出,论文《Real-time Scene Text Detection with Differentiable Binarization》发表于 AAAI 2020;其升级版 DB++(《Real-Time Scene Text Detection with Differentiable Binarization and Adaptive Scale Fusion》)发表于 TPAMI 2022。
传统的分割类文本检测方法通常分两步:先由分割网络输出概率图,再通过固定阈值二值化得到文本区域,最后用形态学操作还原文本框。这一流程中阈值是手工设定的,不可学习,且二值化过程不可导,无法参与端到端训练。DB 算法的核心创新正是把“二值化”本身改造成可微分的操作:
- 网络同时输出概率图(probability map)与阈值图(threshold map);
- 通过可微分的近似阶跃函数(可微分二值化)把两者融合成近似二值图(binary map),训练时梯度可以正常回传;
- 推理阶段阈值图不再需要,仅对概率图做固定阈值切分,因此几乎没有额外推理开销。
在 PaddleOCR 的实现中,这一近似阶跃函数位于 ppocr/modeling/heads/det_db_head.py:
def step_function(self, x, y): return paddle.reciprocal(1 + paddle.exp(-self.k * (x - y)))其中k(默认 50)为放大系数,k越大,该函数越接近真正的阶跃函数。
DB++ 则在 DB 基础上引入**自适应尺度融合(Adaptive Scale Fusion,ASF)**模块,对多尺度特征进行注意力加权融合,同时搭配可变形卷积(DCN)骨干网络,进一步提升弯曲、多尺度文本的检测精度。
2. 公开数据集复现效果
在 ICDAR2015 文本检测公开数据集上,PaddleOCR 官方复现效果如下:
| 模型 | 骨干网络 | 配置文件 | precision | recall | Hmean |
|---|---|---|---|---|---|
| DB | ResNet50_vd | configs/det/det_r50_vd_db.yml | 86.41% | 78.72% | 82.38% |
| DB | MobileNetV3 | configs/det/det_mv3_db.yml | 77.29% | 73.08% | 75.12% |
| DB++ | ResNet50 | configs/det/det_r50_db++_icdar15.yml | 90.89% | 82.66% | 86.58% |
在 TD_TR 文本检测公开数据集上,复现效果如下:
| 模型 | 骨干网络 | 配置文件 | precision | recall | Hmean |
|---|---|---|---|---|---|
| DB++ | ResNet50 | configs/det/det_r50_db++_td_tr.yml | 92.92% | 86.48% | 89.58% |
从上表可以看出:MobileNetV3 骨干的轻量版 DB 在精度略降的情况下大幅压缩计算量,适合移动端部署;而 DB++ 通过 ASF 与 DCN 的加持,在 ICDAR2015 上将 Hmean 从 82.38% 提升到 86.58%。
3. 环境配置与项目准备
训练与推理前需要先配置 PaddleOCR 运行环境并克隆项目代码:
- 运行环境准备请参考 《运行环境准备》;
- 项目代码克隆请参考 《项目克隆》。
4. 配置文件深度解析:DB 与 DB++ 的差异
PaddleOCR 将检测模型模块化为 Backbone(骨干网络)、Neck(特征融合)、Head(检测头)、Loss(损失)、PostProcess(后处理)五大部分,训练不同检测模型只需更换配置文件。下面以 configs/det/det_r50_vd_db.yml(DB)与 configs/det/det_r50_db++_icdar15.yml(DB++)为例逐段讲解。
4.1 Architecture:架构定义
DB 的架构定义:
Architecture: model_type: det algorithm: DB Transform: Backbone: name: ResNet_vd layers: 50 Neck: name: DBFPN out_channels: 256 Head: name: DBHead k: 50DB++ 的架构定义:
Architecture: model_type: det algorithm: DB++ Transform: null Backbone: name: ResNet layers: 50 dcn_stage: [False, True, True, True] # 第 2~4 个 stage 使用可变形卷积 Neck: name: DBFPN out_channels: 256 use_asf: True # 开启自适应尺度融合 ASF Head: name: DBHead k: 50两者关键差异:
- 骨干网络:DB 使用
ResNet_vd;DB++ 使用ResNet并开启dcn_stage,在第 2、3、4 个 stage 使用可变形卷积(DCN)增强几何形变建模能力; - Neck:DB++ 在
DBFPN上增加use_asf: True,启用 ASF 注意力融合模块。ASF 的实现位于 ppocr/modeling/necks/db_fpn.py,它先通过空间注意力(spatial_scale)与通道注意力(channel_scale)为各尺度特征图生成注意力分数,再对p5/p4/p3/p2四层特征加权融合; - Head:两者均使用
DBHead,参数k为可微分二值化的放大系数,默认 50。
DBHead内部包含两个结构相同的子网络:binarize(输出概率图)与thresh(输出阈值图),训练时两者共同前向并融合出近似二值图,详见 ppocr/modeling/heads/det_db_head.py。
4.2 Loss:DB 专用损失函数
DB 与 DB++ 均使用DBLoss,配置如下:
Loss: name: DBLoss balance_loss: true main_loss_type: DiceLoss # DB++ 配置为 BCELoss alpha: 5 beta: 10 ohem_ratio: 3参数含义(对应 ppocr/losses/det_db_loss.py 的实现):
alpha/beta:概率图损失与阈值图损失的加权系数,默认 5 和 10;ohem_ratio:负样本采样比例。训练时通过在线难例挖掘(OHEM)控制正负样本比例,negative_ratio=3表示负样本数量最多为正样本的 3 倍,实现见 ppocr/losses/det_basic_loss.py 中的BalanceLoss;main_loss_type:概率图主损失类型。DB 使用DiceLoss,DB++ 使用BCELoss(通过BalanceLoss包装,同样支持 OHEM)。
DB 的总损失由三部分组成:loss_shrink_maps(概率图,权重alpha)+loss_threshold_maps(阈值图,权重beta)+loss_binary_maps(近似二值图 Dice 损失),见DBLoss.forward的实现。
4.3 Optimizer:优化器配置
DB 使用 Adam 优化器:
Optimizer: name: Adam beta1: 0.9 beta2: 0.999 lr: learning_rate: 0.001 regularizer: name: 'L2' factor: 0DB++ 使用 Momentum + 学习率衰减:
Optimizer: name: Momentum momentum: 0.9 lr: name: DecayLearningRate learning_rate: 0.007 epochs: 1000 factor: 0.9 end_lr: 0 weight_decay: 0.00014.4 PostProcess:后处理参数
PostProcess: name: DBPostProcess thresh: 0.3 # 概率图二值化阈值 box_thresh: 0.7 # 框置信度阈值(DB++ 为 0.6) max_candidates: 1000 # 最大候选框数 unclip_ratio: 1.5 # 外扩比例 det_box_type: 'quad' # 'quad' 或 'poly',DB++ 配置中有该字段对应 ppocr/postprocess/db_postprocess.py 中的DBPostProcess:
thresh:对概率图做阈值切分的固定阈值,对应源码中的segmentation = pred > self.thresh;box_thresh:候选框平均得分阈值,低于该值的候选框被过滤(对应box_score_fast/box_score_slow的得分比较);unclip_ratio:文本框外扩比例,通过pyclipper按area * unclip_ratio / length的距离对多边形做膨胀(unclip方法);max_candidates:最多保留的候选轮廓数;det_box_type:输出框类型,quad输出四点四边形,poly输出多边形(适合弯曲文本)。
4.5 Metric:评估指标
Metric: name: DetMetric main_indicator: hmean检测模型以hmean(F1 分数,precision 与 recall 的调和平均)作为主评估指标。
4.6 Train / Eval:数据与数据增强
DB 的训练数据配置(configs/det/det_r50_vd_db.yml):
Train: dataset: name: SimpleDataSet data_dir: ./train_data/icdar2015/text_localization/ label_file_list: - ./train_data/icdar2015/text_localization/train_icdar2015_label.txt ratio_list: [1.0] transforms: - DecodeImage: img_mode: BGR channel_first: False - DetLabelEncode: - IaaAugment: augmenter_args: - { 'type': Fliplr, 'args': { 'p': 0.5 } } - { 'type': Affine, 'args': { 'rotate': [-10, 10] } } - { 'type': Resize, 'args': { 'size': [0.5, 3] } } - EastRandomCropData: size: [640, 640] max_tries: 50 keep_ratio: true - MakeBorderMap: shrink_ratio: 0.4 thresh_min: 0.3 thresh_max: 0.7 - MakeShrinkMap: shrink_ratio: 0.4 min_text_size: 8 - NormalizeImage: scale: 1./255. mean: [0.485, 0.456, 0.406] std: [0.229, 0.224, 0.225] order: 'hwc' - ToCHWImage: - KeepKeys: keep_keys: ['image', 'threshold_map', 'threshold_mask', 'shrink_map', 'shrink_mask'] loader: shuffle: True drop_last: False batch_size_per_card: 16 num_workers: 4要点说明:
- 标签编码:
DetLabelEncode解析检测标注;MakeShrinkMap根据标注文本多边形生成收缩图(shrink map),MakeBorderMap生成阈值图及其掩码,二者共同构成 DB 训练所需的监督信号; - 数据增强:
IaaAugment提供水平翻转、旋转(-10°~10°)、随机缩放(0.5~3 倍)等增强;EastRandomCropData在 640×640 区域内随机裁剪; - 归一化:
NormalizeImage使用 ImageNet 均值方差归一化,DB 的 mean/std 为[0.485, 0.456, 0.406]/[0.229, 0.224, 0.225];DB++ 使用 SynthText 统计的 mean[0.48109378172549, 0.45752457890196, 0.40787054090196]、std 为 1.0; - Eval 预处理:评估时使用
DetResizeForTest固定测试尺寸(DB 为[736, 1280],DB++ 为[1152, 2048]),loader 的batch_size_per_card必须为 1。
5. 模型训练、评估与预测
DB / DB++ 的训练、评估与预测流程与 PaddleOCR 通用文本检测流程一致,详见 文本检测训练教程。PaddleOCR 对代码做了模块化设计,训练不同检测模型只需更换配置文件,核心命令示例:
# 单卡训练 python3 tools/train.py -c configs/det/det_r50_vd_db.yml # 评估 python3 tools/eval.py -c configs/det/det_r50_vd_db.yml -o Global.pretrained_model=./output/det_r50_vd/best_accuracy # 预测 python3 tools/infer_det.py -c configs/det/det_r50_vd_db.yml -o Global.infer_img=./doc/imgs_en/img_10.jpg训练 DB++ 时只需将-c参数替换为 configs/det/det_r50_db++_icdar15.yml(或 configs/det/det_r50_db++_td_tr.yml),网络结构、损失与后处理会自动按配置切换。
6. 推理部署
6.1 Python 推理
第一步:导出 inference model。将训练保存的模型转换为推理模型,以 ResNet50_vd 骨干、ICDAR2015 英文数据集训练的 DB 模型为例:
python3 tools/export_model.py -c configs/det/det_r50_vd_db.yml -o Global.pretrained_model=./det_r50_vd_db_v2.0_train/best_accuracy Global.save_inference_dir=./inference/det_db第二步:执行检测推理。使用 tools/infer/predict_det.py 进行文本检测:
python3 tools/infer/predict_det.py --image_dir="./doc/imgs_en/img_10.jpg" --det_model_dir="./inference/det_db/" --det_algorithm="DB"可视化文本检测结果默认保存到./inference_results文件夹,结果文件名称前缀为det_res。
注意:ICDAR2015 数据集仅包含 1000 张训练图像,且主要针对英文场景,因此上述模型对中文文本图像的检测效果会比较差。如需中文检测,应使用中文数据集(如 ICDAR2017 MLT、合成中文数据)重新训练或直接使用 PaddleOCR 官方发布的中文检测模型。
6.2 C++ 推理
准备好推理模型后,参考 C++ 推理部署教程 操作即可,PaddleOCR 提供基于 Paddle Inference 的 C++ 端推理方案。
6.3 Serving 服务化部署
准备好推理模型后,参考 Paddle Serving 部署教程 进行服务化部署,支持Python Serving与C++ Serving两种模式,可将 DB 检测模型封装为 HTTP/RPC 服务对外提供能力。
6.4 更多推理部署方式
- Paddle2ONNX 推理:准备好推理模型后,参考 Paddle2ONNX 转换教程 将模型转换为 ONNX 格式,从而在其他推理框架或平台上运行。
7. FAQ
在训练或部署 DB / DB++ 模型时常见问题:
- 检测框不完整或缺失:可适当调大
unclip_ratio(文本框外扩比例)或调低box_thresh(框置信度阈值); - 小文本漏检:调小
thresh或在训练数据中增加小尺度文本样本,也可调整DetResizeForTest的测试尺寸使其更适配输入图像分辨率; - 显存不足:降低
Train.loader.batch_size_per_card,或调小训练时EastRandomCropData的size; - 训练不收敛:确认
pretrained_model指向正确的预训练权重路径,并检查数据集的label_file_list路径与标注格式是否正确。
引用
若在学术工作中使用了本文涉及的算法,请引用以下论文:
@inproceedings{liao2020real, title={Real-time scene text detection with differentiable binarization}, author={Liao, Minghui and Wan, Zhaoyi and Yao, Cong and Chen, Kai and Bai, Xiang}, booktitle={Proceedings of the AAAI Conference on Artificial Intelligence}, volume={34}, number={07}, pages={11474--11481}, year={2020} } @article{liao2022real, title={Real-Time Scene Text Detection with Differentiable Binarization and Adaptive Scale Fusion}, author={Liao, Minghui and Zou, Zhisheng and Wan, Zhaoyi and Yao, Cong and Bai, Xiang}, journal={IEEE Transactions on Pattern Analysis and Machine Intelligence}, year={2022}, publisher={IEEE} }【免费下载链接】PaddleOCR飞桨多语言OCR工具包(实用超轻量OCR系统,支持80+种语言识别,提供数据标注与合成工具,支持服务器、移动端、嵌入式及IoT设备端的训练与部署) Awesome multilingual OCR toolkits based on PaddlePaddle (practical ultra lightweight OCR system, support 80+ languages recognition, provide data annotation and synthesis tools, support training and deployment among server, mobile, embedded and IoT devices)项目地址: https://gitcode.com/paddlepaddle/PaddleOCR
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考