news 2026/10/8 8:12:01

docTR 版面检测模型训练完全指南:基于 references/layout 脚本从零跑通 LW-DETR 训练

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
docTR 版面检测模型训练完全指南:基于 references/layout 脚本从零跑通 LW-DETR 训练
  • 人工智能
  • 深度学习
  • 计算机视觉
  • OCR

【免费下载链接】doctr

docTR (Document Text Recognition) - a seamless, high-performing & accessible library for OCR-related tasks powered by Deep Learning. Ongoing development and maintenance by t2k.

项目地址:https://gitcode.com/gh_mirrors/do/doctr
点击查看免费下载

导读

本文围绕 docTR(Document Text Recognition)官方示例训练脚本 references/layout/train.py 展开,系统讲解如何用 docTR 训练文档版面检测(Layout Detection)模型。读完本文,你将掌握:环境搭建与依赖安装、单卡/多卡(torchrun + DDP)训练命令、设备与混合精度(AMP)配置、数据集的目录结构与labels.json标注规范、checkpoint 与运行元数据的产出逻辑,以及脚本全部命令行参数的含义与默认值,可直接照搬到自己的版面标注数据上完成训练与评估。

一、版面检测在 docTR 中的定位

版面检测的目标是从文档图像中定位并分类出各个版面区域(如标题、正文、表格、图片、页眉、页脚等)。在 docTR 的整体管线中,版面检测通常作为 OCR 的"前置步骤",先找到文档中每个语义区块,再对区块内的文本做识别。

当前仓库中,版面检测由doctr/models/layout/模块实现,核心架构为LW-DETR("LW-DETR: A Transformer Replacement to YOLO for Real-Time Detection",见 doctr/models/layout/lw_detr/base.py 的引用说明)。doctr/models/layout/zoo.py 中定义了可用的预训练架构名单:

架构名说明
lw_detr_sLW-DETR small 版本,默认输入(3, 1024, 1024),带预训练权重(发布在 docTR v1.0.1 release)
lw_detr_mLW-DETR medium 版本,同一输入尺寸,仓库中未提供预训练权重(url: None)

两个架构的默认类别均为 11 类:Caption、Footnote、Formula、List-item、Page-footer、Page-header、Picture、Section-header、Table、Text、Title(见 doctr/models/layout/lw_detr/pytorch.py)。训练时类别列表会由你的数据集动态决定,脚本会以验证集为准并把类别名广播到所有 DDP rank。

二、环境搭建

官方推荐先以可编辑模式安装 docTR,再安装 references 所需的额外依赖(references/layout/README.md):

pip install -e . --upgrade pip install -r references/requirements.txt

references/requirements.txt 的内容如下,除了以可编辑模式安装当前仓库本体(-e .)外,还包括:

  • tqdm:训练进度条(配合 Slack 日志扩展)
  • slack-sdk:Slack 日志推送 SDK
  • wandb>=0.10.31:Weights & Biases 实验跟踪
  • clearml>=1.11.1:ClearML 实验跟踪
  • matplotlib>=3.1.0:样本可视化、LR Finder 曲线绘制

其中wandb、clearml、slack-sdk都是可选能力:脚本在导入阶段会判断环境变量是否配置(见下文 Slack 部分),只有显式使用--wb、--clearml、--push-to-hub等参数时才会真正依赖它们。

三、启动一次训练

最简单的单卡训练命令(references/layout/README.md):

python references/layout/train.py lw_detr_s --train_path path/to/your/train_set --val_path path/to/your/val_set --epochs 5

其中位置参数arch必须是lw_detr_s或lw_detr_m;--train_path与--val_path为必填项,分别指向训练/验证集目录。脚本在 references/layout/train.py 中通过argparse解析全部参数,并默认使用ArgumentDefaultsHelpFormatter,因此随时可以查看:

python references/layout/train.py --help

训练主流程速览

从 references/layout/train.py 的main()可以梳理出一次完整训练的内部流程,理解它有助于定位问题:

  1. 分布式探测:读取WORLD_SIZE环境变量判断是否处于torchrun启动的 DDP 环境;
  2. 设备解析:分布式下每个进程绑定自己的 GPU rank;单进程下由resolve_device自动选择 CUDA → MPS → CPU;
  3. 临时实例取配置:用pretrained=False构建一次模型,取出cfg["mean"]、cfg["std"]用于归一化;
  4. 加载数据集:分别构建LayoutDataset(训练集带增强、验证集不带增强),并对验证集labels.json计算 SHA256 哈希;
  5. 构建模型:layout.__dict__args.arch,支持--resume断点续训;
  6. 训练循环:每 epoch 依次执行fit_one_epoch(训练)+evaluate(验证,输出 loss、mAP@[.5:.95]、AP@[.5]、AP@[.75]),loss 下降时保存 checkpoint,支持--early-stop;
  7. 收尾:按需推送 Hugging Face Hub、销毁进程组。

四、多卡训练(torchrun + DDP)

脚本使用 PyTorch 内置的torchrun作为分布式启动器,由它负责设置LOCAL_RANK、RANK、WORLD_SIZE等环境变量(references/layout/README.md)。分布式模式与单卡模式唯一的额外参数是--backend,用于指定DistributedDataParallel的通信后端,默认nccl(Linux 多卡下通常最快;若你的操作系统不可用可换gloo等)。

常用 torchrun 参数

  • --nproc_per_node=<N>:在本机启动 N 个进程,通常等于要使用的 GPU 数;
  • --nnodes=<M>:参与任务的节点总数,默认 1;
  • --rdzv_backend、--rdzv_endpoint、--rdzv_id:多节点任务的 rendezvous 设置。

GPU 选择

默认会使用所有可见 GPU。若只想用部分卡,必须在运行torchrun之前通过CUDA_VISIBLE_DEVICES限制,例如只使用 CUDA 设备 0 和 2(references/layout/README.md):

CUDA_VISIBLE_DEVICES=0,2 \ torchrun --nproc_per_node=2 references/layout/train.py \ lw_detr_s \ --train_path path/to/train \ --val_path path/to/val \ --epochs 5 \ --backend nccl

DDP 下的实现细节

从源码结构可以推断,脚本在分布式方面做了不少工程化处理(references/layout/ddp_utils.py):

  • ShardSampler:验证集按 rank 做无填充分片(步长采样range(rank, len, num_replicas)),保证每个样本只被验证一次;
  • barrier_download:让 rank 0 先完成 docTR 权重缓存下载,其余 rank 通过 barrier 等待,避免多进程同时下载冲突;
  • sync_val_metric:通过all_reduce汇总各 rank 的验证指标累加器、all_gather_object拼接 GT/预测缓冲区,使summary()反映完整验证集;
  • 类名通过dist.broadcast_object_list从 rank 0 广播到所有 rank;
  • 训练集使用DistributedSampler(shuffle=True, drop_last=True),每个 epoch 需调用sampler.set_epoch(epoch)保证跨 epoch 的 shuffle 随机性。

五、设备选择与混合精度(AMP)

脚本接受统一的--device参数:可以是 CUDA 索引(如0)、cuda:N、mps(Apple Silicon GPU)或cpu。不传时自动按 "CUDA → MPS → CPU" 顺序探测(references/layout/utils.py 的resolve_device实现)。在分布式模式(torchrun)下该参数被忽略,每个进程使用自己的 GPU。

--amp开启自动混合精度,仅支持 CUDA,脚本会在 CPU/MPS 设备上直接抛错(references/layout/train.py)。配合--amp-dtype可切换精度类型(references/layout/utils.py):

  • float16(默认):配合torch.amp.GradScaler做损失缩放;
  • bfloat16:适合 Ampere 及以上架构 GPU。bfloat16 拥有与 float32 相同的指数范围,无需损失缩放,且能规避 float16 在某些 loss 上的溢出问题(references/layout/train.py 的_autocast/_scaler实现)。

两条可直接运行的示例命令(references/layout/README.md):

# Apple Silicon:开启 MPS 回退,少数 MPS 不支持的算子落到 CPU 执行 PYTORCH_ENABLE_MPS_FALLBACK=1 python references/layout/train.py lw_detr_s --train_path path/to/train --val_path path/to/val --epochs 5 --device mps # NVIDIA GPU 使用 bfloat16 混合精度 python references/layout/train.py lw_detr_s --train_path path/to/train --val_path path/to/val --epochs 5 --device 0 --amp --amp-dtype bfloat16

训练循环中的 AMP 路径还包含梯度裁剪clip_grad_norm_(max_norm=1)与scaler.unscale_的配合,保证缩放后的梯度在裁剪前被还原(references/layout/train.py)。

六、Checkpoint 与运行元数据

每次运行会在 checkpoint 旁边生成一个与实验同名的<experiment name>.json元数据文件,且只写入一次(对应整个 run,而非单个 checkpoint)。其中记录的信息(references/layout/README.md、references/layout/utils.py)包括:

  • 架构名称与任务设置:class_names、assume_straight_pages(等于not args.rotation);
  • 本地数据时的数据集哈希:训练集与验证集labels.json的 SHA256;
  • 环境版本:doctr_version、torch_version、git_revision;
  • 完整参数列表:args的逐项 dump,以及learning_rate、backbone_learning_rate、epochs、batch_size、input_size、optimizer、scheduler、rotation、amp等训练配置。

这些信息足以在推理阶段重建模型结构(类别数、是否假设页面摆正)并复现训练过程,而不必记忆训练时的细节。checkpoint 本体由save_checkpoint以<output_dir>/<name>.pt形式保存,保存时机有两个:验证 loss 创新低时自动保存;开启--save-interval-epoch时每个 epoch 额外保存一份_epoch{N}版本(references/layout/utils.py、references/layout/train.py)。

七、数据集格式规范

训练必须同时提供train_path与val_path,每个路径下的目录结构必须严格为"1 个子目录 + 1 个文件"(references/layout/README.md):

├── images │ ├── sample_img_01.png │ ├── sample_img_02.png │ ├── sample_img_03.png │ └── ... └── labels.json

labels.json是一个字典:key 为图像文件名,value 为包含 4 个条目的字典:

字段含义
img_dimensions图像的空间形状(高, 宽)
img_hash图像文件的 SHA256 哈希
polygons定位多边形,由若干 2D 点构成(每个区域 4 个点)
classes每个多边形对应的类别名列表

多边形内部点的顺序无关紧要;点的坐标是(x, y) 绝对坐标。示意如下:

{ "sample_img_01.png": { "img_dimensions": [900, 600], "img_hash": "theimagedumpmyhash", "polygons": [[[x1, y1], [x2, y2], [x3, y3], [x4, y4]], ...], "classes": ["class_name_1", "class_name_2", ...] }, "sample_img_02.png": { "img_dimensions": [900, 600], "img_hash": "thisisahash", "polygons": [[[x1, y1], [x2, y2], [x3, y3], [x4, y4]], ...], "classes": ["class_name_1", "class_name_2", ...] } }

数据加载的源码级校验

数据集由doctr/datasets/layout.py中的LayoutDataset负责加载,其校验规则(doctr/datasets/layout.py)可以作为标注质量的检查清单:

  • labels.json与图像文件必须真实存在,否则抛FileNotFoundError;
  • 每条标注必须包含polygons与classes,缺失抛KeyError;
  • 每个图像的polygons数量必须等于classes数量,否则抛ValueError;
  • 多边形数组形状必须是(N, 4, 2),即每个区域 4 个二维点;use_polygons=False(默认)时会把四边形折叠为最小外接矩形(xmin, ymin, xmax, ymax),即min/max聚合(doctr/datasets/layout.py);
  • 空图像(没有任何区域)是合法的全背景样本;
  • 类别名通过sorted(set(...))统一排序后暴露给模型,因此训练/验证集的类别集合必须一致,脚本会做显式校验(references/layout/train.py)。

此外,--rotation模式下use_polygons变为 True,四边形会被保留为多边形并参与旋转增强与评估,而--eval-straight可让指标评估改用直线框以节省时间与内存(references/layout/train.py)。

八、通过 tqdm 推送 Slack 日志

脚本支持把训练进度直接推送到 Slack 频道(references/layout/README.md)。只需配置两个环境变量:

  • TQDM_SLACK_TOKEN:Slack Bot Token;
  • TQDM_SLACK_CHANNEL:在 Slack 中"右键频道 → Copy → Copy link"可得到形如https://xxxxxx.slack.com/archives/yyyyyyyy的链接,只保留最后一段yyyyyyyy作为频道 ID。

脚本在导入时会检测这两个环境变量,若都存在则从tqdm.contrib.slack导入 Slack 版tqdm,否则回退到tqdm.auto(references/layout/train.py)。主进程中还会 monkey-patchpbar.write,使进度信息直接通过chat_postMessage发送到指定频道(references/layout/train.py)。创建 Slack App 的具体步骤请参考 Slack 官方 Quickstart。

九、实验跟踪:W&B 与 ClearML

脚本内置了两种主流实验跟踪平台(references/layout/train.py):

  • --wb:Weights & Biases,项目名固定为layout-detection,记录训练/验证 loss、学习率与 mAP 指标;
  • --clearml:ClearML,项目名docTR/layout-detection,同样上传 config 与标量指标。

二者可同时开启,由统一的log_at_step在训练与验证循环中调用(每 step 记一次)。注意这两个参数仅在 rank 0 生效,分布式下不会重复记录。

十、高级选项:完整参数清单

脚本通过 references/layout/train.py 的parse_args()暴露了丰富的训练选项,这里给出完整清单与默认值,方便你按需定制:

参数默认值说明
arch(位置参数)必填模型架构:lw_detr_s/lw_detr_m
--backendncclDDP 通信后端
--device自动探测单进程设备:CUDA 索引 /cuda:N/mps/cpu
--output_dir.checkpoint 与最终模型保存目录
--train_path必填训练数据目录
--val_path必填验证数据目录
--nameNone实验名称(默认{arch}_{时间戳})
--epochs10训练轮数
-b, --batch_size2批大小
--save-interval-epoch关每个 epoch 都保存一份 checkpoint
--input_size1024模型输入尺寸,H = W
--lr0.001学习率(Adam / AdamW)
--backbone-lr0.001主干网络(backbone)学习率
--wd, --weight-decay1e-4权重衰减
-j, --workers自动(最多 16)数据加载 worker 数
--resumeNone从 checkpoint 续训
--test-only关只跑验证循环(评估)
--freeze-backbone关冻结主干网络参数做微调
--show-samples关显示未归一化的训练样本
--wb关记录到 Weights & Biases
--clearml关记录到 ClearML
--push-to-hub关训练结束后推送到 Hugging Face Hub(会触发登录)
--pretrained关用预训练参数初始化
--rotation关使用旋转文档训练(多边形标注)
--eval-straight关评估改用直线框,省时间/内存
--optimadamw优化器:adam/adamw
--schedcosine调度器:cosine/onecycle/poly
--amp关自动混合精度(仅 CUDA)
--amp-dtypefloat16AMP 精度:float16/bfloat16
--find-lr关学习率网格搜索(LR Finder)
--early-stop关启用早停
--early-stop-epochs5早停耐心值
--early-stop-delta0.01早停最小改善阈值

几个值得展开的进阶功能

  • LR Finder(--find-lr):以1e-7为起点、1为终点的指数增长方式网格搜索最优学习率,并用指数移动平均绘制"学习率-损失"曲线(references/layout/train.py、references/layout/utils.py)。分布式下仅 rank 0 执行,其余 rank 通过 barrier 等待。
  • 调度器细节:cosine与poly都采用"线性预热 + 主调度"的SequentialLR组合,预热步数取max(1, min(2000, 5% * total_steps));onecycle使用余弦退火策略(references/layout/train.py)。
  • 参数分组优化:build_param_groups将参数按"是否 backbone(feat_extractor.*)"与"是否需权重衰减"分为 4 组,bias、norm、LayerNorm、BN 与 embedding 参数不参与权重衰减,backbone 可用独立学习率(references/layout/utils.py)。
  • 早停:EarlyStopper在验证 loss 不再下降且超出min_delta达到patience个 epoch 时终止训练(references/layout/utils.py)。

十一、数据增强与旋转训练

训练集默认启用一套组合增强(references/layout/train.py),可以直观理解--rotation开关的作用:

  • 图像级增强(img_transforms,随机选一):颜色反转 + 高斯模糊;随机阴影 + 高斯噪声 + 模糊 + 随机灰度;torchvision 的光度扰动;或恒等映射;
  • 样本级增强(图像 + 目标同步变换):随机水平翻转(0.15)、随机裁剪(比例 0.85–1.15、尺度 0.75–1.0)、随机缩放(0.4–0.9),最后统一 Resize 到input_size × input_size并保留纵横比、对称 padding,同时返回 padding mask(版面模型训练必需);
  • --rotation模式下额外加入RandomRotate(90)(旋转 0°/90°/180°/270°,概率 0.5),此时数据以多边形而非直线框参与训练,且use_polygons=True。

验证集不应用增强,只做 Resize + padding mask(--rotation且非--eval-straight时才加入旋转分支)。这与训练/验证同分布但保持验证指标稳定可复现的工程惯例一致。

十二、评估指标与结果解读

验证环节使用 docTR 的ObjectDetectionMetric(doctr/utils/metrics.py),每个 epoch 结束后输出三类指标(references/layout/train.py):

  • mAP@[.5:.95]:IoU 阈值从 0.5 到 0.95 的均值平均精度,综合反映定位质量;
  • AP@[.5]、AP@[.75]:在 IoU 0.5 / 0.75 阈值下的平均精度,分别宽松与严格。

若 GT 或预测为空导致指标未定义,脚本会输出(Undefined metric value, caused by empty GTs or predictions)提示,而不是给出误导性的 0 值。评估以样本数加权汇总 loss,保证结果与 DDP 分片方式无关(references/layout/train.py)。

十三、把训练产物接入推理

训练完成后,除了save_checkpoint保存的*.pt权重,脚本在--push-to-hub时还会通过push_to_hf_hub上传模型。权重与元数据 JSON 配合即可在推理侧重建模型:docTR 的layout_predictor(arch, pretrained=..., assume_straight_pages=...)(doctr/models/layout/zoo.py)支持传入架构名或模型对象,并通过PreProcessor完成归一化与缩放;assume_straight_pages决定后处理输出直线框还是旋转多边形(doctr/models/layout/lw_detr/base.py 中的LWDETRPostProcessor:sigmoid 打分、topk 筛选、置信度阈值 0.5、IoU 阈值 0.5 的逐类 NMS,并按assume_straight_pages折叠为(xmin, ymin, xmax, ymax)或保留(N, 4, 2)多边形)。

从源码结构看,LW-DETR 的后处理针对旋转框做了专门设计:模型输出以 OBB 参数(cx, cy, w, h, sinθ, cosθ)表达,解码时通过cv2.boxPoints恢复四边形并用 ProbIoU(高斯框概率 IoU,参考 "Gaussian Bounding Boxes and Probabilistic IoU")参与训练损失(doctr/models/layout/lw_detr/pytorch.py)。

结语

借助 references/layout 这套官方示例脚本,你可以用统一的命令在单卡与多卡环境下完成 docTR 版面检测模型的数据准备、训练、评估与导出全流程。关键要点可以概括为一句话:数据结构决定训练能否启动,--rotation/--amp/--pretrained等开关决定训练方式,checkpoint 旁的 JSON 元数据保证训练可复现、可重建。动手前先用python references/layout/train.py --help核对参数,再对照第七节的labels.json规范准备数据,即可顺畅跑通整个流程。

  • 人工智能
  • 深度学习
  • 计算机视觉
  • OCR

【免费下载链接】doctr

docTR (Document Text Recognition) - a seamless, high-performing & accessible library for OCR-related tasks powered by Deep Learning. Ongoing development and maintenance by t2k.

项目地址:https://gitcode.com/gh_mirrors/do/doctr
点击查看免费下载

相关推荐

上一篇:微信聊天记录永久保存实战:用 WeChatMsg 把三年对话做成一份可检索的数字档案
下一篇:Video2X 免费 AI 视频放大指南:老片变 4K、30 帧补到 60 帧的完整玩法

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

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

Claude的“外置大脑”:用claude-mem实现跨会话持久记忆

Claude 用久了&#xff0c;大家应该都有同一个磨人的体验&#xff1a;开一个新对话窗口&#xff0c;它就“失忆”了。明明上一个 session 里已经把技术方案、命名规范、部署链路聊得清清楚楚&#xff0c;换一个窗口全部归零&#xff0c;又得从“我们之前讨论过的那件事……”开…

作者头像 李华
网站建设 2026/10/8 8:07:00

编辑器上下文模式实战:用符号树与LSP终结长文件迷失

前阵子我接手一个维护了快七年的老服务&#xff0c;业务逻辑倒不算难&#xff0c;真正的敌人是文件长度。一个核心类 800 多行&#xff0c;方法之间相互调用&#xff0c;我经常滚到屏幕中间就忘了自己是在类里还是已经在某个私有方法里&#xff0c;debug 到一半才发现改动放错了…

作者头像 李华
网站建设 2026/10/8 8:06:37

微信小程序+Java互助学习系统毕设:从技术选型到部署答辩完整指南

简介&#xff1a;这份资源是面向高校计算机相关专业学生与Java初学者的一套互助学习平台毕业设计完整方案&#xff0c;采用微信小程序前端搭配Java后端与MySQL数据库&#xff0c;适合作为毕业设计、课程设计或全栈入门练手项目。压缩包共1306个文件&#xff0c;约25.37MB&#…

作者头像 李华