news 2026/9/13 16:42:43

DiT 文档布局分析实战:基于 Detectron2 的 Mask R-CNN / Cascade Mask R-CNN 推理、训练与评估完全指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
DiT 文档布局分析实战:基于 Detectron2 的 Mask R-CNN / Cascade Mask R-CNN 推理、训练与评估完全指南

DiT 文档布局分析实战:基于 Detectron2 的 Mask R-CNN / Cascade Mask R-CNN 推理、训练与评估完全指南

【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm

本指南以 unilm 仓库中 dit/object_detection 模块为核心,系统讲解如何使用 DiT(Document Image Transformer)作为主干网络,在 Detectron2 框架上完成文档布局分析(Document Layout Analysis)中的目标检测任务,覆盖 PubLayNet 与 ICDAR 2019 cTDaR 两大数据集的推理、数据准备、评估与微调全流程。读者读完本文后,将能够独立复现 DiT-Base / DiT-Large 搭配 Mask R-CNN / Cascade Mask R-CNN 的检测方案,并理解其配置体系与底层实现原理。

1. 模块概述:DiT 的文档布局分析实现

dit/object_detection是 DiT(Document Image Transformer)在文档布局分析任务上的官方实现。该目录基于 Meta 的 Detectron2(Mask R-CNN 与 Cascade Mask R-CNN 的实现来源),将 DiT 预训练权重作为视觉主干(backbone)接入两阶段检测框架,面向两个文档数据集:

  • PubLayNet:大规模文档版面数据集,检测类别为 5 类 ——texttitlelisttablefigure
  • ICDAR 2019 cTDaR:表格检测与识别竞赛数据集,目标类别为table(区分 modern 与 archival 两个子集)。

目录的核心文件布局如下:

路径作用
inference.py单图推理与结果可视化脚本
train_net.py训练 / 评估统一入口
convert_to_coco_format.py将 ICDAR 2019 cTDaR 原始标注转为 COCO 格式
adaptive_binarize.py对 archival 子集做自适应二值化
publaynet_configsPubLayNet 的 Mask R-CNN / Cascade Mask R-CNN 配置
icdar19_configsICDAR 2019 cTDaR 的 Mask R-CNN / Cascade Mask R-CNN 配置
ditodDiT 主干、数据集映射、评估器与训练器扩展实现
publaynet_example.jpeg推理演示示例图片

其中ditod子包是整个方案的“引擎舱”,包含 backbone.py(ViT + FPN 主干)、beit.py 与 deit.py(DiT / BEiT / DEiT / MAE 模型定义)、config.py(ViT 专属配置项)、dataset_mapper.py(DETR 式数据增强)、mytrainer.py(自定义训练器)以及 icdar_evaluation.py(ICDAR 评估器)。

2. 推理:快速体验 DiT 文档布局分析

2.1 使用 inference.py 进行单图推理

官方提供了 Hugging Face Spaces 网页演示,可直接在线体验文档布局分析效果;而在本地,最快的验证方式是运行inference.py脚本。以下命令需在unilm 仓库根目录执行:

python ./dit/object_detection/inference.py \ --image_path ./dit/object_detection/publaynet_example.jpeg \ --output_file_name output.jpg \ --config ./dit/object_detection/publaynet_configs/maskrcnn/maskrcnn_dit_base.yaml \ --opts MODEL.WEIGHTS https://layoutlm.blob.core.windows.net/dit/dit-fts/publaynet_dit-b_mrcnn.pth

务必保证配置(YAML)与 PyTorch 权重匹配。上例使用的是 DiT-Base 主干 + Mask R-CNN 框架在 PubLayNet 上微调后的权重publaynet_dit-b_mrcnn.pth;若改用 DiT-Large 配置,则需替换为对应的dit_large权重,否则加载 checkpoint 时会因网络结构不匹配而失败。

四个命令行参数的含义分别为:

参数说明
--image_path输入图片路径(必填)
--output_file_name可视化结果输出文件名(如output.jpg
--config检测框架配置文件路径
--opts覆盖配置项,格式为KEY VALUE键值对(此处用于指定微调权重)

2.2 推理脚本源码解析

从 inference.py 的源码可以看到推理的完整链路:

  1. 构造配置get_cfg()创建 Detectron2 默认配置,随后调用add_vit_config(cfg)注入MODEL.VIT.*等 DiT 专属配置项(见 config.py),再merge_from_file读取 YAML、merge_from_list应用--opts覆盖;
  2. 设定设备device = "cuda" if torch.cuda.is_available() else "cpu",自动回退到 CPU(CPU 推理速度较慢,仅适合快速验证);
  3. 构建预测器DefaultPredictor(cfg)由 Detectron2 提供,会自动加载MODEL.WEIGHTS指定权重并对输入做ResizeShortestEdge预处理;
  4. 设置类别元数据:根据测试数据集名动态指定类别 —— 若cfg.DATASETS.TEST[0] == 'icdar2019_test'则类别为["table"],否则为["text","title","list","table","figure"](PubLayNet 五类);
  5. 推理与可视化VisualizerColorMode.SEGMENTATION模式绘制预测实例(框 + 掩码 + 类别),最终通过cv2.imwrite保存。

这一流程清晰展示了 Detectron2 “配置驱动”的工程范式:模型结构完全由配置文件决定,脚本只负责组装。

3. 数据集准备

3.1 PubLayNet

PubLayNet 数据集约 96GB,需从官方发布渠道下载publaynet.tar.gz后解压到目录PATH-to-PubLayNet。解压后执行:

ln -s PATH-to-PubLayNet publaynet_data

软链接名称必须为publaynet_data。其原因在 train_net.py 中写死:脚本通过register_coco_instances注册数据集时,硬编码了"./publaynet_data/train.json""./publaynet_data/train"等相对路径。因此在 unilm 仓库根目录下创建该软链接,程序才能访问到数据。

3.2 ICDAR 2019 cTDaR

ICDAR 2019 cTDaR 数据约 4GB,下载后假设仓库路径名为PATH-to-ICDARrepo。首先将原始数据转换为 COCO 格式:

python convert_to_coco_format.py --root_dir=PATH-to-ICDARrepo --target_dir=PATH-toICDAR

处理后的数据位于PATH-to-ICDAR。接着对archival 子集执行自适应二值化(现代印刷体 modern 子集无需处理):

cp -r PATH-to-ICDAR/trackA_archival PATH-to-ICDAR/at_trackA_archival python adaptive_binarize.py --root_dir PATH-to-ICDAR/at_trackA_archival

二值化后的 archival 子集保存在PATH-to-ICDAR/at_trackA_archival。随后根据要评估/微调的子集,在仓库根目录建立data软链接:

ln -s PATH-to-ICDAR/trackA_modern data # 评估 modern 子集 # 或 ln -s PATH-to-ICDAR/at_trackA_archival data # 评估 archival 子集

与 PubLayNet 同理,train_net.py 中注册 ICDAR 数据时使用的是"data/train.json""data/test.json"等相对路径,因此软链接必须命名为data且建立在当前工作目录。

3.3 数据预处理脚本源码解读

convert_to_coco_format.py 的核心逻辑是将 ICDAR 的 XML 标注解析为 COCO JSON:类别固定为单一tablecategories: [{"id": 1, "name": "table"}]);从 XML 的<table/Coords>节点读取表格四角点,计算 segmentation 多边形与bbox(取[x0, y0, x3-x0, y3-y0]);文件名前缀cTDaR_t0对应trackA_archivalcTDaR_t1对应trackA_modern。该脚本还内置clean_img()函数,用于统一.JPG.TIFF.png等图片格式为.jpg

adaptive_binarize.py 使用 OpenCV 的cv2.adaptiveThresholdADAPTIVE_THRESH_GAUSSIAN_C,blockSize=45,C=11)对灰度图做高斯自适应阈值二值化,再转回三通道 BGR 覆写原图,以提升档案扫描件的表格边界检测效果。

4. 评估:验证微调后的检测性能

评估使用 train_net.py 的--eval-only模式。配置文件位于icdar19_configspublaynet_configs两个目录。

示例 1:评估 PubLayNet 上微调的 DiT-Base + Mask R-CNN:

python train_net.py --config-file publaynet_configs/maskrcnn/maskrcnn_dit_base.yaml --eval-only --num-gpus 8 MODEL.WEIGHTS <finetuned_checkpoint_file_path or link> OUTPUT_DIR <your_output_dir>

示例 2:评估 ICDAR 2019 cTDaR archival 子集上微调的 DiT-Large + Cascade Mask R-CNN(需先将PATH-to-ICDAR/at_trackA_archival软链接为data):

python train_net.py --config-file icdar19_configs/cascade/cascade_dit_large.yaml --eval-only --num-gpus 8 MODEL.WEIGHTS <finetuned_checkpoint_file_path or link> OUTPUT_DIR <your_output_dir>

4.1 ICDAR 2019 测量工具的 Bug 修复

重要提示:官方在将 ICDAR2019 测量工具集成进代码时,修复了原工具中的一个 bug。如果你使用外部 ICDAR 测量工具(ctdar_measurement_tool)自行计算评估分数,请将evaluate.py中按扩展名过滤 ground-truth 文件的代码修改如下(原代码for file in gt_file_lst: ... gt_file_lst.remove(file)在遍历过程中删除列表元素会导致漏删或越界):

... # print(each_file) # for file in gt_file_lst: # if file.split(".") != "xml": # gt_file_lst.remove(file) # # print(gt_file_lst) # Comment the code above and add the code below for i in range(len(gt_file_lst) - 1, -1, -1): if gt_file_lst[i].split(".")[-1] != "xml": del gt_file_lst[i] if len(gt_file_lst) > 0: ...

即改为从后往前倒序遍历并删除,避免“边遍历边删除”造成的元素遗漏问题。

4.2 评估器实现

仓库自带的评估分流逻辑位于 mytrainer.py 的build_evaluator:数据集名包含icdar时使用自定义ICDAREvaluator(见 icdar_evaluation.py,其内部集成了修复后的 ICDAR 测量逻辑),其余情况使用 Detectron2 标准COCOEvaluator

5. 训练:微调 DiT 主干

以下两条命令展示了如何使用 DiT 主干 + Mask R-CNN / Cascade Mask R-CNN 在8 张 32GB NVIDIA V100 GPU上进行微调。

示例 1:PubLayNet 上微调 DiT-Base + Cascade Mask R-CNN:

python train_net.py --config-file publaynet_configs/cascade/cascade_dit_base.yaml --num-gpus 8 MODEL.WEIGHTS <DiT-Base_file_path or link> OUTPUT_DIR <your_output_dir>

示例 2:ICDAR 2019 cTDaR modern 子集上微调 DiT-Large + Mask R-CNN:

python train_net.py --config-file icdar19_configs/markrcnn/maskrcnn_dit_large.yaml --num-gpus 8 MODEL.WEIGHTS <DiT-Large_file_path or link> OUTPUT_DIR <your_output_dir>

微调时MODEL.WEIGHTS传入的是DiT 自监督预训练权重(如dit-base-224-p16-500k-62d53a.pthdit-large-224-p16-500k-d7a2fb.pth),由 maskrcnn_dit_base.yaml 等配置的MODEL.WEIGHTS字段指定;命令行的MODEL.WEIGHTS覆盖则用于指定已微调 checkpoint(配合--eval-only)或替换预训练权重来源。更详细的 Detectron2 用法可参考其官方文档。

5.1 训练入口与数据集注册

train_net.py 的main()首先通过register_coco_instances注册四个数据集(publaynet_train/valicdar2019_train/test),随后setup(args)完成配置合并与冻结。它复用了 Detectron2 的launch()分布式启动器,支持--num-gpus--num-machines--machine-rank--dist-url等标准参数,并额外提供--debug参数(内部使用 debugpy 在 0.0.0.0:9310 等待调试器附加)。

5.2 自定义训练器与优化策略

MyTrainer 继承自 Detectron2 的TrainerBase,其中几个关键设计点:

  • 数据加载:当cfg.AUG.DETR=True时使用自定义 DetrDatasetMapper,启用 DETR 风格增强 —— 以 50% 概率插入ResizeShortestEdge([400,500,600]) + RandomCrop(absolute_range, (384,600))裁剪序列;
  • AMP 混合精度cfg.SOLVER.AMP.ENABLED=True时训练循环自动切换为AMPTrainer
  • 优化器:支持 SGD / AdamW,并为 backbone 参数提供BACKBONE_MULTIPLIER学习率缩放;启用full_model梯度裁剪时,会在 step 前对整个模型参数执行clip_grad_norm_(见build_optimizer);
  • 调度器:使用WarmupCosineLR,配置了WARMUP_FACTOR=0.01WARMUP_ITERS等;
  • 钩子:内置IterationTimerLRSchedulerPreciseBNPeriodicCheckpointerEvalHookPeriodicWriter等训练钩子,并在训练期间按TEST.EVAL_PERIOD自动做周期性评估。

6. 配置文件逐项解析

配置文件采用 Detectron2 的 YACS 继承体系:子配置通过_BASE_: "../Base-RCNN-FPN.yaml"继承公共配置,再按数据集与模型规模覆盖差异项。

6.1 公共配置 Base-RCNN-FPN.yaml

publaynet_configs/Base-RCNN-FPN.yaml 定义了检测框架的公共结构:

配置项说明
MODEL.META_ARCHITECTUREGeneralizedRCNN标准两阶段检测架构
MODEL.MASK_ONTrue启用实例分割分支
MODEL.PIXEL_MEAN/STD[123.675, 116.280, 103.530]/[58.395, 57.120, 57.375]图像归一化参数(DiT 配置会覆盖为 127.5 系)
MODEL.BACKBONE.NAMEbuild_vit_fpn_backbone注册到BACKBONE_REGISTRY的 ViT+FPN 主干
MODEL.VIT.OUT_FEATURES["layer3","layer5","layer7","layer11"]从 DiT 提取的多尺度特征层
MODEL.VIT.DROP_PATH0.1随机深度(Stochastic Depth)丢弃率
MODEL.VIT.IMG_SIZE[224,224]预训练输入分辨率
MODEL.VIT.POS_TYPEabs绝对位置编码
MODEL.FPN.IN_FEATURESOUT_FEATURES相同FPN 输入特征
MODEL.ROI_HEADS.NUM_CLASSES5PubLayNet 五类
SOLVER.BASE_LR0.0004基础学习率
SOLVER.IMS_PER_BATCH32全局 batch size
INPUT.CROPabsolute_range (384,600)DETR 式随机裁剪
INPUT.MIN_SIZE_TRAIN(480,512,...,800)短边随机缩放范围
AUG.DETRTrue启用 DETR 数据增强
SEED42随机种子

6.2 VIT 配置注入(add_vit_config)

所有配置文件都必须先经 config.py 的add_vit_config(cfg)注入MODEL.VIT.*默认值,否则会出现“配置项不存在”错误。其注册的默认值包括:

  • MODEL.VIT.NAME(默认""):主干模型名,可选dit_base_patch16dit_large_patch16beit_base_patch16beit_large_patch16deit_base_patch16mae_base_patch16
  • MODEL.VIT.OUT_FEATURES(默认["layer3","layer5","layer7","layer11"]):输出哪些 Transformer 层的特征;
  • MODEL.VIT.IMG_SIZE(默认[224,224]);
  • MODEL.VIT.POS_TYPE(默认"shared_rel"):位置编码类型,可取值abs/shared_rel/rel
  • MODEL.VIT.DROP_PATH(默认0.);
  • MODEL.VIT.MODEL_KWARGS(默认"{}"):透传给模型构造函数的额外参数;
  • SOLVER.OPTIMIZER(默认"ADAMW")、SOLVER.BACKBONE_MULTIPLIER(默认1.0);
  • AUG.DETR(默认False):是否启用 DETR 数据增强。

6.3 各数据集与模型规模的配置差异

PubLayNet Mask R-CNN(DiT-Base)—— maskrcnn_dit_base.yaml:覆盖PIXEL_MEAN/STD[127.5, 127.5, 127.5](与 DiT 预训练归一化一致);MODEL.VIT.NAME: "dit_base_patch16"WARMUP_ITERS: 1000IMS_PER_BATCH: 16MAX_ITER: 60000CHECKPOINT_PERIOD: 2000TEST.EVAL_PERIOD: 2000

PubLayNet Cascade Mask R-CNN(DiT-Base)—— cascade_dit_base.yaml:在 Mask R-CNN 基础上将ROI_HEADS.NAME改为CascadeROIHeadsROI_BOX_HEAD.CLS_AGNOSTIC_BBOX_REG: True(类别无关的框回归)、RPN.POST_NMS_TOPK_TRAIN: 2000

ICDAR 2019(DiT-Large)—— maskrcnn_dit_large.yaml 与 cascade_dit_large.yaml 的共同差异:MODEL.VIT.NAME: "dit_large_patch16"OUT_FEATURESFPN.IN_FEATURES切换为["layer7","layer11","layer15","layer23"](24 层 DiT-Large 的深层特征);DROP_PATH: 0.2;学习率降至BASE_LR: 0.00005IMS_PER_BATCH: 16;checkpoint / 评估周期缩短为1000

7. DiT 主干网络实现原理

7.1 VIT_Backbone 与 FPN 的组装

backbone.py 中的VIT_Backbone负责将 ViT 模型包装为 Detectron2 的Backbone,其_out_feature_strides按模型规模区分:

  • Base 系列dit_base_patch16等,12 层):layer3→stride 4layer5→8layer7→16layer11→32
  • Large 系列dit_large_patch16beit_large_patch16,24 层):layer7→4layer11→8layer15→16layer23→32

build_vit_fpn_backbone(注册为build_vit_fpn_backbone)在VIT_Backbone之上叠加 Detectron2 标准FPNtop_block使用LastLevelMaxPool生成p6层,最终形成 P2–P6 特征金字塔供 RPN 与 ROI Heads 使用。

7.2 DiT 模型结构与多尺度特征输出

DiT 的模型定义位于 beit.py 的BEiT类。从源码看,dit_base_patch16dit_large_patch16的差异主要体现在:

  • embed_dim / 深度 / 头数:Base 为 768 / 12 层 / 12 头,Large 为 1024 / 24 层 / 16 头;
  • LayerScale 初值:Base 为init_values=0.1,Large 为init_values=1e-5(残差分支乘以可学习缩放向量,见Block中的gamma_1/gamma_2);
  • 两者均使用qkv_bias=True、patch_size=16、mlp_ratio=4。

forward_features逐层前向,当层号命中out_indices时,将 token 序列重排回二维特征图(去掉cls_tokenreshape(B, C, Hp, Wp)),最后经过四个轻量 FPN 头生成多尺度输出:patch16 场景下fpn1为两层ConvTranspose2d上采样(stride 4)、fpn2为单层上采样(stride 8)、fpn3Identity(stride 16)、fpn4MaxPool2d(stride 32),对应 backbone 中声明的主干 stride 映射。use_checkpoint=True时,各 Block 通过torch.utils.checkpoint做激活重计算以节省显存,这也是 8×V100 32GB 能跑 DiT-Large 的重要原因。

7.3 位置编码与推理灵活性

BEiT支持三种位置编码(对应配置POS_TYPE):

  • abs:可学习的绝对位置编码(use_abs_pos_emb=True);
  • shared_rel:跨层共享的相对位置偏置(RelativePositionBias,在Attention中加到注意力分数上);
  • rel:每层独立的窗口内相对位置偏置。

其中RelativePositionBias实现了 bicubic 插值,当推理分辨率与预训练[224,224]不一致时,可自动将位置偏置表插值到新的窗口尺寸(见 beit.py 中training_window_size != window_size的分支)。配合 deit.py 中的interpolate_pos_encoding,使得 DiT 主干能处理文档检测所需的任意分辨率输入。

8. 注意事项与常见问题

  1. 配置与权重必须匹配:DiT-Base 配置配 DiT-Base 预训练/微调权重,DiT-Large 同理;混用会导致 checkpoint 加载失败或精度异常。
  2. 数据集软链接命名:必须在运行命令的目录下创建publaynet_data(PubLayNet)与data(ICDAR)软链接,因为 train_net.py 中数据集注册路径是硬编码的相对路径。
  3. archival 子集需要二值化:仅对 ICDAR 的trackA_archival执行adaptive_binarize.py,modern 子集直接使用原始扫描图。
  4. ICDAR 测量工具 Bug:若使用第三方 ctdar_measurement_tool 复算分数,必须按上文修复evaluate.py中 gt 文件过滤逻辑。
  5. 运行环境:本模块依赖 Detectron2(Mask R-CNN / Cascade Mask R-CNN 实现)与 timm 库;训练推荐 8×32GB V100,配置中的IMS_PER_BATCHBASE_LRMAX_ITER以 8 卡为基准。
  6. 归一化参数:DiT 配置将PIXEL_MEAN/STD覆盖为[127.5, 127.5, 127.5],与 DiT 预训练(mean=0.5、std=0.5,等价于归一化到 [-1,1])保持一致,切勿沿用 Detectron2 默认的 ImageNet 归一化。

9. 引用与致谢

如果本仓库对您的研究或工程有所帮助,请引用 DiT 论文:

@misc{li2022dit, title={DiT: Self-supervised Pre-training for Document Image Transformer}, author={Junlong Li and Yiheng Xu and Tengchao Lv and Lei Cui and Cha Zhang and Furu Wei}, year={2022}, eprint={2203.02378}, archivePrefix={arXiv}, primaryClass={cs.CV} }

特别感谢 Detectron2 项目提供的 Mask R-CNN 与 Cascade Mask R-CNN 实现,以及 DETR / DINO / timm / BEiT 等开源工作为 DiT 检测分支带来的工程基础。

【免费下载链接】unilmLarge-scale Self-supervised Pre-training Across Tasks, Languages, and Modalities项目地址: https://gitcode.com/GitHub_Trending/un/unilm

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

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

CAN总线故障排查:先查物理层再抓报文

1. 大多数人查CAN总线故障&#xff0c;第一步就错了——不是看报文&#xff0c;而是先“听”物理层你有没有遇到过这样的场景&#xff1a;整车报“网关通信超时”&#xff0c;诊断仪读出一串UDS故障码&#xff08;比如U0100、U0121&#xff09;&#xff0c;工程师立刻打开CANoe…

作者头像 李华
网站建设 2026/9/13 16:38:04

Sa-Token 前后端分离鉴权实战:无 Cookie 模式下 Token 的下发、存储与提交

Sa-Token 前后端分离鉴权实战&#xff1a;无 Cookie 模式下 Token 的下发、存储与提交 【免费下载链接】Sa-Token ✨ 开源、免费、一站式 Java 权限认证框架&#xff0c;让鉴权变得简单、优雅&#xff01;—— 登录认证、权限认证、分布式 Session 会话、微服务网关鉴权、SSO 单…

作者头像 李华
网站建设 2026/9/13 16:36:26

电商导购返利小程序实战:淘宝京东拼多多联盟API对接与uniapp开发

简介&#xff1a;首席省钱赚钱专家v1.9.18小程序源码&#xff0c;面向个人创业者、电商运营与小程序开发者&#xff0c;基于拼多多优惠商品接口&#xff0c;实现购物返利、推广分销、团队奖励等典型电商小程序功能&#xff0c;帮助快速搭建“自购省钱分享赚钱”的应用场景。资源…

作者头像 李华
网站建设 2026/9/13 16:36:17

基于JSP的银行预约管理系统:从Servlet原理到并发排错实战

简介&#xff1a;这是一份基于JSPSQLServerTomcat技术栈的银行预约管理系统毕业设计源码包&#xff0c;面向Java Web方向的毕业生或需要快速搭建预约类管理系统的开发者&#xff0c;解决银行业务预约、客户信息管理、后台审核等环节的一体化实现问题。资源共499个文件&#xff…

作者头像 李华