PaddleOCR 中 DRRG 任意形状文本检测算法:图推理原理、配置详解与训练实践
【免费下载链接】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
DRRG(Deep Relational Reasoning Graph Network)是一种基于图卷积推理的任意形状文本检测算法,由 Zhang 等人发表于 CVPR 2020,其核心思想是把弯曲文本实例建模为"文本组件构成的图",再通过 GCN 推理组件间的连接关系,从而输出贴合曲线文字的任意形状边界。本文以 PaddleOCR 官方 DRRG 算法文档 为主线,结合仓库中完整的配置、源码与后处理实现,系统讲解 DRRG 的算法原理、det_r50_drrg_ctw.yml全量参数、数据标签生成、损失函数、后处理流水线以及基于 CTW1500 的训练/评估/预测实操命令,帮助读者既能在 PaddleOCR 中一键复现该算法,也能深入理解其内部机制。
一、算法简介与复现结果
DRRG 的论文信息如下:
Deep Relational Reasoning Graph Network for Arbitrary Shape Text Detection Zhang, Shi-Xue and Zhu, Xiaobin and Hou, Jie-Bo and Liu, Chang and Yang, Chun and Wang, Hongfa and Yin, Xu-Cheng CVPR, 2020
与 DBNet、EAST 等直接回归文本框/分割图的方案不同,DRRG 从"文本由一系列有序组件(component)构成"这一观察出发:先把文本区域分解为若干近似四边形的小组件,把每个组件当作图中的一个节点,然后用图卷积网络(GCN)判断节点之间是否存在"同属一个文本实例"的连接边,最后把连接成簇的组件合并成任意形状的多边形边界。这一思路天然适合弯曲、倾斜、旋转的文本行。
在 PaddleOCR 中,DRRG 使用 CTW1500 文本检测公开数据集训练,复现效果如下(来源:算法文档):
| 模型 | 骨干网络 | 配置文件 | Precision | Recall | Hmean |
|---|---|---|---|---|---|
| DRRG | ResNet50_vd | configs/det/det_r50_drrg_ctw.yml | 89.92% | 80.91% | 85.18% |
二、模型结构与源码实现
DRRG 在 PaddleOCR 中遵循"Backbone + Neck + Head"的模块化结构,配置文件Architecture一节定义如下:
Architecture: model_type: det algorithm: DRRG Transform: Backbone: name: ResNet_vd layers: 50 Neck: name: FPN_UNet in_channels: [256, 512, 1024, 2048] out_channels: 32 Head: name: DRRGHead in_channels: 32 text_region_thr: 0.3 center_region_thr: 0.42.1 骨干与特征融合 Neck
- Backbone:
ResNet_vd、layers: 50,即 ResNet50_vd,训练时通常加载ResNet50_vd_ssld_pretrained.pdparams预训练权重(见配置文件Global.pretrained_model)。 - Neck:
FPN_UNet,输入四级特征图通道数[256, 512, 1024, 2048],输出 32 通道融合特征。其实现位于 ppocr/modeling/necks/fpn_unet.py,内部由 4 个UpBlock(1x1 卷积 + 3x3 卷积 + 转置卷积)和一个up4转置卷积组成,逐级上采样并与编码器特征融合,最终输出通道数为 32 的统一尺度特征图。
2.2 DRRGHead:从特征图到预测图
DRRGHead 是 DRRG 的核心头部,它用一个 1x1 卷积out_conv把 32 通道特征映射为6 通道预测图,各通道含义依次为:
pred_text_region—— 文本区域得分图;pred_center_region—— 文本中心区域得分图;pred_sin_map—— 中心区域像素到文本顶/底边方向的 sin(θ);pred_cos_map—— 中心区域像素到文本顶/底边方向的 cos(θ);pred_top_height_map—— 中心区域像素到上边线的距离;pred_bot_height_map—— 中心区域像素到下边线的距离。
前向过程中,DRRGHead会把输入特征与 6 通道预测图在通道维拼接(paddle.concat([inputs, pred_maps], axis=1)),得到用于构建图节点内容的特征。
在训练阶段,head 通过LocalGraphs(ppocr/modeling/heads/local_graph.py)基于 GT 组件属性构建局部图,送入GCN预测连接关系;在推理阶段,head 则通过ProposalLocalGraphs(ppocr/modeling/heads/proposal_local_graph.py)从预测图上自动提议文本组件并构建局部图。DRRGHead的完整默认超参数如下:
| 参数 | 默认值 | 含义 |
|---|---|---|
k_at_hops | (8, 4) | 一跳/两跳邻居数量,决定局部图扩展范围 |
num_adjacent_linkages | 3 | 邻接矩阵中每个节点连接的近邻数 |
node_geo_feat_len | 120 | 节点几何特征嵌入长度 |
pooling_scale | 1.0 | RoIAlignRotated 采样尺度 |
pooling_output_size | (4, 3) | 旋转 RoI 池化输出尺寸 |
text_region_thr | 0.2 | 文本区域阈值(配置文件覆盖为 0.3) |
center_region_thr | 0.2 | 中心区域阈值(配置文件覆盖为 0.4) |
local_graph_thr | 0.7 | 训练时局部图去重 IoU 阈值 |
2.3 GCN 图推理模块
GCN 是 DRRG 的"推理大脑",结构为:BatchNorm1D → 4 层 GraphConv(512→256→128→64)→ 分类头(Linear(64,32) + PReLU + Linear(32,2))。其中GraphConv采用均值聚合(MeanAggregator,即bmm(A, features)),把邻接矩阵 A 与节点特征相乘得到聚合特征,再与原始特征拼接后做线性变换加 ReLU。
GCN 的输入节点特征由两部分拼接而成(见LocalGraphs.__call__):
- 内容特征:对每个文本组件,用旋转 RoI Align(
RoIAlignRotated,池化尺寸(4, 3))从"输入特征 + 6 通道预测图"拼接后的特征图中抽取,展平后得到4*3*(32+6) = 456维向量; - 几何特征:把组件的
(x, y, h, w, cos, sin)六元几何属性通过正弦/余弦位置编码嵌入到node_geo_feat_len=120维(feature_embedding,见 local_graph.py)。
节点特征维度为456 + 120 = 576,即GCN(feat_len=576)。GCN 输出每个"候选边"的二分类得分,判断两个组件是否属于同一文本实例。训练时LocalGraphs.generate_local_graphs还通过局部图 IoU 去重(local_graph_thr)减少冗余局部图,并基于 GT 标签生成连接关系监督信号(gt_linkage)。
三、配置文件全量解读(det_r50_drrg_ctw.yml)
configs/det/det_r50_drrg_ctw.yml 是 DRRG 在 CTW1500 上的完整训练配置,各节参数说明如下。
3.1 Global 全局配置
Global: use_gpu: true epoch_num: 1200 log_smooth_window: 20 print_batch_step: 5 save_model_dir: ./output/det_r50_drrg_ctw/ save_epoch_step: 100 # evaluation is run every 1260 iterations eval_batch_step: [37800, 1260] cal_metric_during_train: False pretrained_model: ./pretrain_models/ResNet50_vd_ssld_pretrained.pdparams checkpoints: save_inference_dir: use_visualdl: False infer_img: doc/imgs_en/img_10.jpg save_res_path: ./output/det_drrg/predicts_drrg.txtepoch_num: 1200:总训练轮数较大,配合衰减学习率长周期训练;eval_batch_step: [37800, 1260]:前 37800 次迭代不评估,之后每 1260 次迭代评估一次;pretrained_model:ResNet50_vd 的 SSLD 预训练权重路径,训练前需手动下载放置;infer_img/save_res_path:单图预测的输入图片与结果保存路径。
3.2 Optimizer 优化器
Optimizer: name: Momentum momentum: 0.9 lr: name: DecayLearningRate learning_rate: 0.028 epochs: 1200 factor: 0.9 end_lr: 0.0000001 weight_decay: 0.0001使用动量 0.9 的 Momentum 优化器,初始学习率 0.028,采用DecayLearningRate衰减策略:每经过一个 epoch 学习率乘以factor: 0.9,下限为end_lr: 0.0000001,权重衰减 0.0001。
3.3 PostProcess 后处理
PostProcess: name: DRRGPostprocess link_thr: 0.8link_thr: 0.8是图传播阶段判断"两个组件是否相连"的边得分阈值,是影响最终检测精度的关键超参数(实现见 ppocr/postprocess/drrg_postprocess.py)。
3.4 Metric 评估指标
Metric: name: DetFCEMetric main_indicator: hmeanDRRG 与 FCE(Fourier Contour Embedding)等任意形状检测算法一样,采用 FCE 评估协议(DetFCEMetric),以hmean作为主指标,对应论文表格中的 Precision / Recall / Hmean。
3.5 Train 训练数据流水线
Train: dataset: name: SimpleDataSet data_dir: ./train_data/ctw1500/imgs/ label_file_list: - ./train_data/ctw1500/imgs/training.txt transforms: - DecodeImage: # load image img_mode: BGR channel_first: False ignore_orientation: True - DetLabelEncode: # Class handling label - ColorJitter: brightness: 0.12549019607843137 saturation: 0.5 - RandomScaling: - RandomCropFlip: crop_ratio: 0.5 - RandomCropPolyInstances: crop_ratio: 0.8 min_side_ratio: 0.3 - RandomRotatePolyInstances: rotate_ratio: 0.5 max_angle: 60 pad_with_fixed_color: False - SquareResizePad: target_size: 800 pad_ratio: 0.6 - IaaAugment: augmenter_args: - { 'type': Fliplr, 'args': { 'p': 0.5 } } - DRRGTargets: - 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', 'gt_text_mask', 'gt_center_region_mask', 'gt_mask', 'gt_top_height_map', 'gt_bot_height_map', 'gt_sin_map', 'gt_cos_map', 'gt_comp_attribs'] # dataloader will return list in this order loader: shuffle: True drop_last: False batch_size_per_card: 4 num_workers: 8训练数据采用SimpleDataSet读取 CTW1500 的图片目录与标签文件。关键点:
- 数据增强组合覆盖色彩抖动(
ColorJitter)、随机缩放、随机裁剪翻转、多边形随机裁剪(RandomCropPolyInstances)、最大 60° 的随机旋转(RandomRotatePolyInstances)、SquareResizePad(缩放到短边 800、pad 比例 0.6)以及水平翻转(IaaAugment); DRRGTargets是 DRRG 专属的标签生成算子(见下文第四节),一次前向中直接产出 8 个训练目标,因此KeepKeys中列出了全部 8 个键:image、gt_text_mask、gt_center_region_mask、gt_mask、gt_top_height_map、gt_bot_height_map、gt_sin_map、gt_cos_map、gt_comp_attribs;- loader 配置
batch_size_per_card: 4、num_workers: 8。
3.6 Eval 评估数据流水线
Eval: dataset: name: SimpleDataSet data_dir: ./train_data/ctw1500/imgs/ label_file_list: - ./train_data/ctw1500/imgs/test.txt transforms: - DecodeImage: # load image img_mode: BGR channel_first: False ignore_orientation: True - DetLabelEncode: # Class handling label - DetResizeForTest: limit_type: 'min' limit_side_len: 640 - NormalizeImage: scale: 1./255. mean: [0.485, 0.456, 0.406] std: [0.229, 0.224, 0.225] order: 'hwc' - Pad: - ToCHWImage: - KeepKeys: keep_keys: ['image', 'shape', 'polys', 'ignore_tags'] loader: shuffle: False drop_last: False batch_size_per_card: 1 # must be 1 num_workers: 2评估阶段不做任何数据增强,仅DetResizeForTest(限制短边为 640)与Pad对齐尺寸;loader 注释明确要求batch_size_per_card: 1(must be 1),因为 DRRG 推理过程逐图构建局部图,不支持 batch 并行。
四、DRRG 专属标签生成:DRRGTargets
训练 DRRG 需要一组特殊的多通道监督信号,均由 ppocr/data/imaug/drrg_targets.py 中的DRRGTargets在数据加载阶段实时生成(对应配置中的- DRRGTargets:,无参数即全部使用默认值)。其generate_targets会产出 7 类 GT:
gt_text_mask:文本区域掩膜(对所有标注多边形fillPoly填充为 1);gt_mask:有效区域掩膜(被ignore_tags标记的多边形区域置 0,其余为 1);gt_center_region_mask:中心区域掩膜,由顶/底边线按center_region_shrink_ratio=0.3收缩得到;gt_top_height_map/gt_bot_height_map:中心区域每个像素到上/下边线的距离图;gt_sin_map/gt_cos_map:中心区域每个像素处"上点-下点"方向向量的 sin/cos 值;gt_comp_attribs:文本组件属性,形状为(num_max_comps, 8),每行是(num_comps, x, y, h, w, cos, sin, comp_label)。
组件生成过程中,DRRGTargets 会先对每个文本实例的边线进行重采样(resample_step=8.0)、按comp_w_h_ratio=0.3与min_width/max_width=(8.0, 24.0)约束生成近似四边形组件,然后经lanms四边形 NMS(text_comp_nms_thr=0.25)去重,再做随机扰动(jitter_comp_attribs,jitter_level=0.2)以增强鲁棒性;当组件数少于num_min_comps=9时,还会在非文本区域随机采样伪组件补齐。这些 GT 组件属性正是训练阶段LocalGraphs建图与 GCN 监督信号(组件是否同属一个实例)的直接来源。
五、损失函数:六项联合监督
DRRGLoss 将 6 通道预测与 GCN 输出联合起来,总损失为:
loss = loss_text + loss_center + loss_height + loss_sin + loss_cos + loss_gcn各项含义与实现方式:
| 损失项 | 类型 | 说明 |
|---|---|---|
loss_text | 平衡二值交叉熵(Balanced BCE) | 监督文本区域图pred_text_region,负样本按ohem_ratio=3.0取 top-k 困难样本,缓解正负样本不均衡 |
loss_center | BCE | 监督中心区域图pred_center_region,正样本除以文本区域均值、负样本除以非文本区域均值后加权(负样本权重 0.5) |
loss_height | Smooth L1(对数缩放) | 监督pred_top/bot_height_map,以log(gt_height+1)为权重聚焦高文本区域,仅在中心区域内计算 |
loss_sin/loss_cos | Smooth L1 | 监督方向图,预测的 sin/cos 先按sqrt(1/(sin²+cos²))归一化到单位圆上 |
loss_gcn | 交叉熵(CrossEntropy) | 监督 GCN 输出的组件连接二分类,GT 由"两组件是否属于同一文本实例"构成 |
训练时 head 返回(pred_maps, (gcn_pred, gt_labels))二元组,DRRGLoss从labels[1:8]中取出 7 个 GT 张量并逐项计算,最终forward返回包含loss及各分量明细的字典,便于训练日志观测(对应print_batch_step: 5的打印)。
六、后处理流水线:从边到任意形状边界
推理时DRRGHead.single_test返回三元组(edges, scores, text_comps),随后由 DRRGPostprocess(配置link_thr: 0.8)完成以下步骤:
graph_propagation:把边按组件中心距离过滤(edge_len_thr=50.0之外置 0 分),去重合并重复边得分,构建无向图节点(Node类,带links集合);connected_components:以link_thr为阈值做连通分量聚类——得分低于阈值的边被剪断,得到若干组件簇;clusters2labels:为每个组件分配簇标签;remove_single:删除孤立单组件簇(抑制误检);comps2boundaries:对每个簇,用min_connect_path求组件中心点的最短连接路径,排序后取上下边线均值生成 top/bot 两条边线,再用fix_corner补全首尾拐角,最终输出2k+1维的任意形状边界点序列(末位为簇平均得分);resize_boundary:按shape_list中的缩放因子把边界还原到原图尺寸。
该后处理对"任意形状"输出至关重要:它把离散的四边形组件通过图聚类与路径规划重新组织成一条贴合弯曲文本的连续多边形边界。
七、环境准备、数据下载与训练/评估/预测
7.1 环境与数据
- 环境配置:参考 《运行环境准备》 安装 PaddlePaddle 与 PaddleOCR 依赖,参考 《项目克隆》 克隆仓库;
- 数据集:CTW1500 的下载说明见 ocr_datasets。按配置文件约定,图片放在
./train_data/ctw1500/imgs/,训练/测试标签文件分别为training.txt与test.txt; - 预训练权重:把
ResNet50_vd_ssld_pretrained.pdparams放入./pretrain_models/,与配置Global.pretrained_model对应。
7.2 训练
PaddleOCR 对代码进行了模块化,训练不同的检测模型只需更换配置文件。基于 文本检测训练教程 的通用命令,DRRG 单卡训练为:
python3 tools/train.py -c configs/det/det_r50_drrg_ctw.yml \ -o Global.pretrained_model=./pretrain_models/ResNet50_vd_ssld_pretrained.pdparams断点续训(指定Global.checkpoints):
python3 tools/train.py -c configs/det/det_r50_drrg_ctw.yml \ -o Global.checkpoints=./your/trained/model多卡分布式训练:
python3 -m paddle.distributed.launch --gpus '0,1,2,3' \ tools/train.py -c configs/det/det_r50_drrg_ctw.yml7.3 评估
python3 tools/eval.py -c configs/det/det_r50_drrg_ctw.yml \ -o Global.checkpoints=./output/det_r50_drrg_ctw/best_accuracy评估采用DetFCEMetric,日志中关注hmean主指标(复现目标为 85.18%)。
7.4 单图预测
python3 tools/infer_det.py -c configs/det/det_r50_drrg_ctw.yml \ -o Global.infer_img="./doc/imgs_en/img_10.jpg" \ Global.pretrained_model="./output/det_r50_drrg_ctw/best_accuracy"也可通过-o Global.infer_img传入图片目录批量预测。
八、推理部署支持情况(重要限制)
官方算法文档明确标注了 DRRG 的部署支持范围:
| 部署方式 | 支持情况 | 原因 |
|---|---|---|
| Python 推理 | 不支持(动态图转静态图) | 模型前向过程中需要多次将张量转换为 Numpy 数据参与运算(局部图构建、lanms NMS、邻接矩阵归一化等),Paddle 动转静机制暂无法覆盖 |
| C++ 推理 | 不支持 | — |
| Serving 服务化部署 | 不支持 | — |
| 更多推理部署 | 不支持 | — |
因此 DRRG 目前主要用于学术复现与 Python 动态图场景下的训练/评估/预测,生产级服务化部署建议改用 DBNet 等支持完整导出链路的检测算法。从源码结构看,推理路径中ProposalLocalGraphs依赖cv2、lanms等 NumPy 生态算子逐组件处理,这也与文档所述限制相互印证。
九、FAQ
- Q:为什么 DRRG 训练需要专门的
DRRGTargets算子?A:因为 GCN 的训练监督信号(组件属性与连接关系标签)必须由 GT 多边形实时生成,无法像 DBNet 那样只依赖简单的二值掩膜;这也是KeepKeys中目标键数量远多于普通检测算法的原因。 - Q:
link_thr对结果影响大吗?A:大。它决定图传播阶段边的保留强度,阈值过高会把一个文本实例切成多段(Recall 下降),过低则会粘连相邻文本行(Precision 下降),建议围绕 0.8 做小范围网格搜索。 - Q:评估时为什么 batch 必须为 1?A:推理建图按单图进行,
batch_size_per_card: 1 # must be 1是配置中明确的硬性约束,见 det_r50_drrg_ctw.yml 的Eval.loader。
引用
@inproceedings{zhang2020deep, title={Deep relational reasoning graph network for arbitrary shape text detection}, author={Zhang, Shi-Xue and Zhu, Xiaobin and Hou, Jie-Bo and Liu, Chang and Yang, Chun and Wang, Hongfa and Yin, Xu-Cheng}, booktitle={Proceedings of the IEEE/CVF Conference on Computer Vision and Pattern Recognition}, pages={9699--9708}, year={2020} }【免费下载链接】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),仅供参考