PaddleOCR 关键信息抽取算法 SDMGR 实战:双模态图推理的原理、配置与训练评估预测全解析
【免费下载链接】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 中基于 Spatial Dual-Modality Graph Reasoning(SDMGR) 论文实现的关键信息抽取(KIE)算法展开,系统讲解其"视觉 + 文本"双模态图推理的算法原理、配套源码实现、wildreceipt 数据集上的完整训练/评估/预测流程,并逐项拆解其核心配置文件。读完本文,你将能够独立完成 SDMGR 模型从数据准备、配置修改到训练评估与可视化的全链路实操,并理解其底层图神经网络的设计细节。
1. 算法简介:SDMGR 解决什么问题
SDMGR(Spatial Dual-Modality Graph Reasoning for Key Information Extraction)是一种关键信息抽取算法,其核心任务是将文档中每个检测到的文本区域(textline)分类为预定义的语义类别,例如订单 ID、发票号码、金额等。与纯文本序列标注方案不同,SDMGR 显式建模文本之间的空间关系(如上下、左右、包含等),将整张票据/文档构造成一张图:文本区域作为节点(node),区域间的空间关系作为边(edge),从而把关键信息抽取转化为图上的节点分类与边分类问题。
论文信息如下:
Spatial Dual-Modality Graph Reasoning for Key Information Extraction
Hongbin Sun, Zhanghui Kuang, Xiaoyu Yue, Chenhao Lin, Wayne Zhang
2021
在 wildreceipt 发票公开数据集上,PaddleOCR 对该算法的复现效果如下:
| 模型 | 骨干网络 | 配置文件 | Hmean | 下载链接 |
|---|---|---|---|---|
| SDMGR | VGG6(UNet 变体) | configs/kie/sdmgr/kie_unet_sdmgr.yml | 86.70% | 官方提供训练模型 kie_vgg16.tar(推理模型待发布) |
在 PaddleOCR 仓库中,SDMGR 的完整实现分布在以下源码文件中,构成了"数据预处理 → 骨干网络 → 图推理头 → 损失 → 评估指标"的完整链路:
- 图推理头:ppocr/modeling/heads/kie_sdmgr_head.py(
SDMGRHead、GNNLayer、Block) - 骨干网络:ppocr/modeling/backbones/kie_unet_sdmgr.py(
Kie_backbone) - 损失函数:ppocr/losses/kie_sdmgr_loss.py(
SDMGRLoss) - 评估指标:ppocr/metrics/kie_metric.py(
KIEMetric) - 预测脚本:tools/infer_kie.py
2. 算法原理与源码级解析
2.1 整体架构:Backbone + SDMGRHead
从 configs/kie/sdmgr/kie_unet_sdmgr.yml 的Architecture段可以看到,SDMGR 的网络结构分为骨干网络与头部两部分:
Architecture: model_type: kie algorithm: SDMGR Transform: Backbone: name: Kie_backbone Head: name: SDMGRHead其中Kie_backbone(定义于 ppocr/modeling/backbones/kie_unet_sdmgr.py)是一个基于 UNet 结构的编码器-解码器网络:编码器由"卷积 + BatchNorm + ReLU + 池化"堆叠而成,逐级提取图像的多尺度视觉特征;解码器通过上采样与跳跃连接恢复分辨率。该骨干的作用是从整张票据图像中提取视觉特征图(visual feature),供后续头部与文本节点特征进行融合。
2.2 节点与边的构造:SDMGRHead 的前向流程
SDMGRHead(ppocr/modeling/heads/kie_sdmgr_head.py)的构造参数如下:
| 参数 | 默认值 | 含义 |
|---|---|---|
in_channels | 必填 | 骨干网络输出的通道数 |
num_chars | 92 | 字符字典大小,用于文本序列嵌入 |
visual_dim | 16 | 视觉 ROI 特征的维度 |
fusion_dim | 1024 | 双模态融合模块的中间维度 |
node_input | 32 | 字符嵌入的维度 |
node_embed | 256 | 节点(文本区域)的嵌入维度 |
edge_input | 5 | 空间关系(边)的原始特征维度 |
edge_embed | 256 | 边的嵌入维度 |
num_gnn | 2 | 图神经网络(GNN)层数 |
num_classes | 26 | 预定义语义类别数量 |
bidirectional | False | 是否使用双向 LSTM 编码文本 |
前向过程forward(self, input, targets)接收(relations, texts, x)三元组:
- 文本节点编码:对每个文本区域的字符索引序列做
nn.Embedding嵌入,再送入单层 LSTM(nn.LSTM),取最后一个有效字符位置对应的隐状态作为该文本区域的节点特征(node_embed维)。 - 视觉特征融合:若存在骨干输出的视觉特征
x,则通过多模态融合模块Block将视觉 ROI 特征与文本节点特征融合(self.fusion([x, nodes])),这正是"双模态"的体现。 - 边编码:将每对文本区域之间的空间关系向量
relations(5 维)通过nn.Linear映射为边嵌入embed_edges,并做 L2 归一化。 - GNN 推理:将节点与边送入堆叠的
num_gnn层GNNLayer进行消息传递与聚合。 - 输出:
self.node_cls(nodes)输出节点类别 logits,self.edge_cls(cat_nodes)输出边类别 logits(2 类,表示关系是否成立)。
2.3 图推理核心:GNNLayer
GNNLayer(同文件内)实现了单层图卷积的聚合逻辑:
- 将每个样本内的节点两两拼接(
paddle.concat([expand(nodes, ...), expand(nodes, ...)], -1))构造出num² × (node_dim*2)的节点对特征,再与边特征拼接后过in_fc线性层与 ReLU; - 通过
coef_fc计算注意力系数,并使用softmax(-eye(num)*1e9 + coefs)屏蔽自环(对角线置为极小值),实现基于注意力权重的邻居聚合; - 聚合结果经
out_fc与 ReLU 后作为残差加到原节点特征上(nodes += relu(out_fc(...))),形成"节点更新 + 残差连接"的图卷积单元。
2.4 双模态融合模块:Block
Block是一个借鉴多模态紧凑双线性池化(MCB)思路的融合模块:两个输入分支分别经过线性层映射到高维空间后,按chunks分块,每块通过rank秩的 Hadamard 积与按秩求和(m = m0(x0_c) * m1(x1_c); z = paddle.sum(m, 1))实现紧凑双线性特征交互,最后再经过正负 ReLU 开方(sqrt(relu(z)) - sqrt(relu(-z)))与归一化输出。该模块用于将骨干提取的视觉特征与 LSTM 编码的文本节点特征深度融合。
2.5 损失函数:SDMGRLoss
SDMGRLoss(ppocr/losses/kie_sdmgr_loss.py)采用节点分类与边分类双分支交叉熵:
loss_node = nn.CrossEntropyLoss(ignore_index=0):节点类别预测损失,索引为 0 的类别被忽略;loss_edge = nn.CrossEntropyLoss(ignore_index=-1):边(空间关系)预测损失,-1 表示无有效关系的位置被忽略;- 最终损失
loss = node_weight * loss_node + edge_weight * loss_edge(默认node_weight=1.0、edge_weight=1.0)。
同时该损失在forward中还会计算节点与边的 Top-1 准确率(acc_node、acc_edge),供训练日志观察。从pre_process的实现可以看出,每个样本的真实标签gts是一个num × (num+1)的矩阵:第一列为节点类别,其余列为该节点与其他节点的关系类别,tag记录每个样本的真实节点数与标签长度。
2.6 评估指标:KIEMetric
KIEMetric(ppocr/metrics/kie_metric.py)实现了文档原论文约定的评估协议:计算混淆矩阵后,忽略掉 26 个类别中 13 个"其他/忽略"类别(如ignores = [0, 2, 4, ..., 24, 25]所列索引),仅对有效类别计算逐类 F1 并取平均作为hmean。配置文件中Metric.main_indicator: hmean即指定该值为最终衡量指标。
3. 环境配置与数据准备
3.1 环境与项目准备
请先参考 《运行环境准备》 配置 PaddleOCR 运行环境,再参考 《项目克隆》 克隆项目代码(中文版环境说明见 environment.md)。
3.2 下载 wildreceipt 数据集
SDMGR 的训练与测试数据来自 wildreceipt 数据集(票据类文档,包含文本行、类别标签与空间关系标注),通过如下命令下载并解压:
wget https://paddleocr.bj.bcebos.com/ppstructure/dataset/wildreceipt.tar && tar xf wildreceipt.tar解压完成后,将数据集软链到PaddleOCR/train_data目录下:
cd PaddleOCR/ && mkdir train_data && cd train_data ln -s ../../wildreceipt ./数据就绪后,目录结构应满足配置文件中的默认路径约定(train_data/wildreceipt/下包含wildreceipt_train.txt、wildreceipt_test.txt、dict.txt、class_list.txt等文件)。
4. 配置文件详解:kie_unet_sdmgr.yml
训练、评估与预测统一使用 configs/kie/sdmgr/kie_unet_sdmgr.yml。下面逐段拆解其关键参数。
4.1 Global 全局配置
Global: use_gpu: True epoch_num: 60 log_smooth_window: 20 print_batch_step: 50 save_model_dir: ./output/kie_5/ save_epoch_step: 50 eval_batch_step: [ 0, 80 ] # 每 80 个 iter 评估一次 load_static_weights: False cal_metric_during_train: False pretrained_model: checkpoints: save_inference_dir: use_visualdl: False class_path: &class_path ./train_data/wildreceipt/class_list.txt infer_img: ./train_data/wildreceipt/1.txt save_res_path: ./output/sdmgr_kie/predicts_kie.txt img_scale: [ 1024, 512 ]要点说明:
epoch_num: 60为总训练轮数;eval_batch_step: [0, 80]表示从第 0 个迭代开始每 80 个迭代执行一次评估;class_path指向类别名称文件(class_list.txt),其行号即类别索引,训练、预测、结果可视化均依赖该映射(tools/infer_kie.py中的read_class_list会逐行读取生成idx -> class_name字典);infer_img指定预测阶段输入的文本文件(每行存储图片路径与 OCR 标注信息的 JSON);save_res_path为预测结果文本文件路径,可视化图片默认保存于其所在目录下的kie_results/子目录;img_scale: [1024, 512]用于KieResize变换,控制输入图像缩放的尺寸。
4.2 Architecture / Loss / Optimizer
Architecture: model_type: kie algorithm: SDMGR Backbone: name: Kie_backbone Head: name: SDMGRHead Loss: name: SDMGRLoss Optimizer: name: Adam beta1: 0.9 beta2: 0.999 lr: name: Piecewise learning_rate: 0.001 decay_epochs: [ 60, 80, 100] values: [ 0.001, 0.0001, 0.00001] warmup_epoch: 2 regularizer: name: 'L2' factor: 0.00005 PostProcess: name: None Metric: name: KIEMetric main_indicator: hmean要点说明:
- 优化器采用Adam(
beta1=0.9、beta2=0.999),学习率采用Piecewise 分段衰减:初始 0.001,分别在 epoch 60、80、100 处衰减为 0.0001、0.00001,并带有 2 个 epoch 的 warmup;L2 权重衰减系数为 0.00005; PostProcess.name: None表示该算法不设置独立后处理模块;Metric指定KIEMetric,以hmean作为主指标。
4.3 Train 训练数据流
Train: dataset: name: SimpleDataSet data_dir: ./train_data/wildreceipt/ label_file_list: [ './train_data/wildreceipt/wildreceipt_train.txt' ] ratio_list: [ 1.0 ] transforms: - DecodeImage: # 加载图像 img_mode: RGB channel_first: False - NormalizeImage: scale: 1 mean: [ 123.675, 116.28, 103.53 ] std: [ 58.395, 57.12, 57.375 ] order: 'hwc' - KieLabelEncode: # 标签编码:节点类别、边关系、文本序列 character_dict_path: ./train_data/wildreceipt/dict.txt class_path: *class_path - KieResize: - ToCHWImage: - KeepKeys: keep_keys: [ 'image', 'relations', 'texts', 'points', 'labels', 'tag', 'shape'] loader: shuffle: True drop_last: False batch_size_per_card: 4 num_workers: 4要点说明:
KieLabelEncode(实现在 ppocr/data/imaug/label_ops.py)负责解析训练标签:dict.txt为字符字典(决定num_chars),class_path为类别文件(决定num_classes);其输出relations(空间关系)、texts(文本字符序列)、points(文本行四点坐标)、labels(节点/边标签矩阵)、tag(每个样本的有效数量信息)正是SDMGRHead与SDMGRLoss的输入;KieResize(ppocr/data/imaug/operators.py)按Global.img_scale对图像与坐标同步缩放;KeepKeys中列出的键顺序即 dataloader 返回列表的顺序,训练与评估阶段的键集合不同(评估阶段额外包含ori_image、ori_boxes,用于结果可视化与评估)。
4.4 Eval 评估数据流
Eval: dataset: name: SimpleDataSet data_dir: ./train_data/wildreceipt label_file_list: - ./train_data/wildreceipt/wildreceipt_test.txt transforms: - DecodeImage: img_mode: RGB channel_first: False - KieLabelEncode: character_dict_path: ./train_data/wildreceipt/dict.txt - KieResize: - NormalizeImage: scale: 1 mean: [ 123.675, 116.28, 103.53 ] std: [ 58.395, 57.12, 57.375 ] order: 'hwc' - ToCHWImage: - KeepKeys: keep_keys: [ 'image', 'relations', 'texts', 'points', 'labels', 'tag', 'ori_image', 'ori_boxes', 'shape'] loader: shuffle: False drop_last: False batch_size_per_card: 1 # 评估时 batch size 必须为 1 num_workers: 4注意:评估阶段的batch_size_per_card必须设置为 1,因为KIEMetric与标签预处理依赖逐样本处理(batch[4].squeeze(0)、tag解析等逻辑假定 batch 内只有一个样本)。
5. 模型训练、评估与预测
5.1 模型训练
配置文件默认训练数据路径为train_data/wildreceipt,数据准备好后执行:
python3 tools/train.py -c configs/kie/sdmgr/kie_unet_sdmgr.yml -o Global.save_model_dir=./output/kie/-o Global.save_model_dir=./output/kie/通过命令行覆盖配置项,将模型保存目录指定为./output/kie/。训练过程中的节点/边准确率、总损失(loss、loss_node、loss_edge)会按print_batch_step: 50的频率打印。
5.2 模型评估
执行下面的命令对训练好的模型进行评估(Global.checkpoints指向保存的最佳模型):
python3 tools/eval.py -c configs/kie/sdmgr/kie_unet_sdmgr.yml -o Global.checkpoints=./output/kie/best_accuracy输出信息示例如下:
[2022/08/10 05:22:23] ppocr INFO: metric eval *************** [2022/08/10 05:22:23] ppocr INFO: hmean:0.8670120239257812 [2022/08/10 05:22:23] ppocr INFO: fps:10.18816520530961其中hmean即文档 §2.6 中KIEMetric忽略非目标类别后计算的平均 F1(0.867 与表格中 86.70% 的复现效果一致),fps为评估吞吐。
5.3 模型预测与结果可视化
SDMGR 的预测由专用脚本 tools/infer_kie.py 完成,与常规 OCR 推理不同,预测时需要预先加载一个存储"图片路径 + OCR 标注信息"的文本文件,通过Global.infer_img指定:
python3 tools/infer_kie.py -c configs/kie/sdmgr/kie_unet_sdmgr.yml -o Global.checkpoints=kie_vgg16/best_accuracy Global.infer_img=./train_data/wildreceipt/1.txt说明:原文档此命令中的配置文件路径写作
configs/kie/kie_unet_sdmgr.yml,但当前仓库中该文件实际位于configs/kie/sdmgr/kie_unet_sdmgr.yml,请以仓库实际路径为准;Global.checkpoints指向官方提供的kie_vgg16预训练模型。
infer_kie.py的执行流程(可从源码确认):
- 通过
read_class_list(class_path)读取class_list.txt构建类别索引映射; - 逐行读取
Global.infer_img指定的文本文件,每行格式为图片相对路径\t标签JSON,其中标签 JSON 数组的每个元素包含该文本区域的transcription(识别文本)与points(四点坐标); - 构建模型并加载权重(
build_model+load_model),执行前向得到节点与边预测; draw_kie_result将预测类别与置信度绘制到图像上——左侧为原图叠加检测框,右侧为标注了"类别(置信度)"的预测图,并保存到save_res_path所在目录的kie_results/子目录(默认./output/sdmgr_kie/kie_results/);write_kie_result将每条文本行的预测结果(label、transcription、score、points)以 JSON 数组形式写入save_res_path(默认./output/sdmgr_kie/predicts_kie.txt),并按预测类别排序输出。
可视化结果示例如下:
从图中可以看到,票据中的"订单号、日期、金额、名称"等文本区域被逐一标注为对应语义类别并附带置信度,直观体现了 SDMGR 将整张票据建模为图、对每个节点(文本区域)做分类的能力。
6. 推理部署支持情况与 FAQ
6.1 推理部署
截至本文撰写时(以当前仓库 docs/version2.x 文档为准),SDMGR 算法的常规推理部署支持情况如下:
- Python 推理:暂不支持(预测请使用 tools/infer_kie.py 脚本);
- C++ 推理部署:暂不支持;
- Serving 服务化部署:暂不支持;
- 更多推理部署(如移动端/其他框架):暂不支持。
该算法的使用范围当前主要面向科研复现与训练/评估流程,正式生产部署前请关注官方后续版本对 SDMGR 推理支持的更新。
6.2 FAQ
本节为占位章节,目前无额外高频问题记录;实际使用中如遇到数据格式相关问题,建议优先核对train_data/wildreceipt/下wildreceipt_train.txt、wildreceipt_test.txt的标签格式是否与KieLabelEncode的解析约定一致。
7. 引用
如需在论文或报告中引用 SDMGR 算法,可使用以下 BibTeX:
@misc{sun2021spatial, title={Spatial Dual-Modality Graph Reasoning for Key Information Extraction}, author={Hongbin Sun and Zhanghui Kuang and Xiaoyu Yue and Chenhao Lin and Wayne Zhang}, year={2021}, eprint={2103.14470}, archivePrefix={arXiv}, primaryClass={cs.CV} }此外,本仓库中还有更多关键信息抽取算法的实现与文档可供参考,例如基于 LayoutLM 的 algorithm_kie_layoutxlm.en.md 与基于 Vi-LayoutXLM 的 algorithm_kie_vi_layoutxlm.en.md,以及对应的中文版文档,可作为 KIE 技术选型与对比研究的延伸阅读。
【免费下载链接】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),仅供参考