news 2026/9/17 23:53:13

在 OOTDiffusion 中用好 detectron2 模型:构建、Checkpoint 加载与输入输出格式全解析

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
在 OOTDiffusion 中用好 detectron2 模型:构建、Checkpoint 加载与输入输出格式全解析

在 OOTDiffusion 中用好 detectron2 模型:构建、Checkpoint 加载与输入输出格式全解析

【免费下载链接】OOTDiffusion[AAAI 2025] Official implementation of "OOTDiffusion: Outfitting Fusion based Latent Diffusion for Controllable Virtual Try-on"项目地址: https://gitcode.com/GitHub_Trending/oo/OOTDiffusion

本指南以 OOTDiffusion 仓库内置的 detectron2 框架文档为核心,系统讲解如何通过build_model构建模型、借助DetectionCheckpointer加载与保存.pth/.pkl权重、按统一的list[dict]输入输出协议调用模型,以及如何部分执行模型以获取中间张量。读完本篇,你将掌握 detectron2 模型从构建、加载到推理的完整调用链,并能把同样的模式套用到 OOTDiffusion 的人体解析(humanparsing)预处理模块上。

一、模型从何而来:build_model与构建函数体系

在 detectron2 中,模型(及其子模型)统一由一组构建函数产出,最常用的是build_modelbuild_backbonebuild_roi_heads等。核心入口如下:

from detectron2.modeling import build_model model = build_model(cfg) # returns a torch.nn.Module

这里传入的cfg是一个已冻结的CfgNode配置对象(见 config/defaults.py),它决定了模型的架构细节:cfg.MODEL.BACKBONE.NAME指定骨干网络、cfg.MODEL.ROI_HEADS.NUM_CLASSES指定类别数、cfg.MODEL.MASK_ON决定是否输出 Mask 分支等。build_model只负责构建模型结构并用随机参数填充,并不会加载任何预训练权重——加载是 Checkpoint 的职责(见下一节)。

在 detectron2/modeling/init.py 中,build_model的实际实现是先根据cfg.MODEL.META_ARCHITECTUREMETA_ARCH_REGISTRY中取出对应的元架构类(如GeneralizedRCNNSemanticSegmentor),再实例化。这种"注册表 + 配置名"的机制意味着你不需要修改框架源码,就能替换任意内部组件,详见 write-models.md。

二、加载与保存 Checkpoint:DetectionCheckpointer实战

构建出的模型参数是随机的,必须从已有权重恢复。detectron2 提供了专门的检查点类:

from detectron2.checkpoint import DetectionCheckpointer # 将权重文件加载进模型 DetectionCheckpointer(model).load(file_path) # 保存到 output/model_999.pth checkpointer = DetectionCheckpointer(model, save_dir="output") checkpointer.save("model_999")

DetectionCheckpointer的实现位于 detectron2/checkpoint/detection_checkpoint.py。从源码(L26-L45)可以看出它支持两类权重格式:

  • .pth文件:PyTorch 原生格式,走 fvcore 的Checkpointer._load_file标准加载流程;若文件内没有model顶层键,会被自动包装为{"model": loaded}
  • .pkl文件:针对 detectron2 model zoo 及旧版 Caffe2/Detectron1 权重的兼容加载。源码会区分两种情况:若数据含model__author__键,则视为 detectron2 model zoo 格式直接使用;否则视为 Caffe2/Detectron1 格式,剥离*_momentum键并设置matching_heuristics=True,随后通过align_and_update_state_dicts按名称匹配启发式完成权重对齐转换(见 c2_model_loading.py)。另外,DetectionCheckpointer默认只在主进程写入磁盘(save_to_disk=is_main_process),这是多卡训练时避免重复写盘的设计。

可移植性提示DetectionCheckpointer只负责"结构无关"的加载,加载后的模型与 checkpoint 中的类别数、Mask 分支配置必须一致。因此在 OOTDiffusion 的人体解析场景中,配置里MODEL.ROI_HEADS.NUM_CLASSESMODEL.MASK_ON必须与权重训练时保持一致(下文第五节会给出具体配置)。

此外,.pth文件也可直接用torch.load/torch.save任意操作,.pkl文件则可用pickle.dump/pickle.load处理,便于做权重剪枝、格式转换等定制。

三、调用模型:训练态与推理态的统一接口

detectron2 中模型的调用约定非常简单:outputs = model(inputs),其中inputslist[dict]——一个 dict 对应一张图像,dict 的键取决于模型类型以及当前处于训练还是评估模式。

训练模式:必须在EventStorage下运行

训练时所有模型都要求在EventStorage上下文内执行,训练统计(各 loss 项)会被写入存储,供日志与可视化使用:

from detectron2.utils.events import EventStorage with EventStorage() as storage: losses = model(inputs) # 返回 dict[str -> ScalarTensor]

推理模式:DefaultPredictor一行封装

如果只想用现成模型做简单推理,DefaultPredictor 是官方推荐的封装,它内置了模型加载、图像预处理,并直接面向单张图像(而非 batch)工作:

from detectron2.engine import DefaultPredictor predictor = DefaultPredictor(cfg) outputs = predictor(image) # image 为 BGR 格式的 np.ndarray

DefaultPredictor内部会完成三件事:按cfg.MODEL.WEIGHTSDetectionCheckpointer加载权重、把图像按cfg.INPUT系列配置(如MIN_SIZE_TESTMAX_SIZE_TEST)缩放到模型输入尺寸、再以推理模式执行模型。

在仓库的 demo/predictor.py 中,VisualizationDemo正是基于DefaultPredictor构建的:run_on_image先把 BGR 图像送入self.predictor(image)拿到predictionsdict,再交给Visualizer绘制实例掩码/语义分割/全景分割可视化结果;而AsyncPredictor(同文件 L132-L220)则用多进程在多个 GPU 上异步推理,通过有界任务队列(task_queue = mp.Queue(maxsize=num_workers * 3))和序号排序保证输出顺序,专门用于加速视频流的可视化吞吐。命令行入口见 demo/demo.py,它把 config 文件、--confidence-threshold等参数合并进cfg后冻结(cfg.freeze())。

四、模型输入格式:list[dict]标准协议

内置模型统一接收list[dict],每个 dict 对应一张图,可能包含以下键:

类型与说明
"image"Tensor,形状(C, H, W)。通道含义由cfg.INPUT.FORMAT决定(如RGB/BGR);归一化在模型内部用cfg.MODEL.PIXEL_MEANcfg.MODEL.PIXEL_STD完成,无需在外部预处理
"instances"Instances 对象(训练时的标注),包含字段:gt_boxes(N 个框的Boxes)、gt_classes(长整型标签向量,取值[0, num_categories))、gt_masksPolygonMasksBitMasks,N 个实例掩码)、gt_keypoints(N 组关键点)
"proposals"Instances对象(仅 Fast R-CNN 风格模型使用),含proposal_boxes(P 个候选框)与objectness_logits(P 个得分向量)
"height","width"期望的输出分辨率,允许与输入image的尺寸不同。例如输入是缩放后的图,但希望输出恢复到原始分辨率时,在 dict 中带上这两个键,模型就会按该分辨率产出结果,比事后插值更高效、更准确
"sem_seg"Tensor[int],形状(H, W),语义分割真值,类别标签从 0 开始

与数据加载器的衔接

默认 DatasetMapper 的输出正是一个符合上述格式的 dict:它对每张图完成读取、resize、随机翻转/裁剪等增强,并组装出instances字段。数据加载器在 batch 之后拼成list[dict],正是内置模型可直接消费的输入。这意味着只要你产出的 dict 符合上表协议,就能绕过 DataLoader 直接喂给模型——这是第五节中"部分执行模型"以及自定义流水线的关键前提。

五、模型输出格式:训练 loss 与推理结果

  • 训练模式:内置模型输出dict[str -> ScalarTensor],键为各 loss 名称(如loss_clsloss_box_regloss_mask),可配合EventStorage记录训练曲线。
  • 推理模式:内置模型输出list[dict],每张图一个 dict,按任务类型可能包含以下字段:
字段类型与说明
"instances"Instances 对象,含pred_boxes(N 个检测框)、scores(N 个得分)、pred_classes(N 个类别标签,[0, num_categories))、pred_masks(形状(N, H, W)的实例掩码 Tensor)、pred_keypoints(形状(N, num_keypoint, 3),最后一维为(x, y, score)且 score > 0)
"sem_seg"Tensor,形状(num_categories, H, W),语义分割预测
"proposals"Instances对象,含proposal_boxes(N 个框)与objectness_logits(N 个得分)
"panoptic_seg"元组(Tensor, list[dict])。Tensor 形状(H, W),每个元素是像素所属的 segment id;每个 dict 描述一个 segment:id(段 id)、isthing(thing 还是 stuff)、category_id(thing 类或 stuff 类对应的类别 id)

六、部分执行模型:获取中间张量

模型内部通常有成百上千个中间张量,官方没有提供"取第 n 个中间结果"的通用 API。文档给出两条可行路径:

路径一:重写(子)模型。参考 write-models.md,通过注册机制改写某个模型组件(如某个 head),使其行为与原组件一致但额外返回你需要的输出。

路径二:部分执行forward()正常构建模型,但不用model(inputs)整体调用,而是按内部模块逐段执行。文档中的示例是"取 Mask head 之前的 mask 特征":

images = ImageList.from_tensors(...) # 预处理后的输入张量 model = build_model(cfg) features = model.backbone(images.tensor) # 骨干网络输出多尺度特征 proposals, _ = model.proposal_generator(images, features) # RPN 生成候选 instances = model.roi_heads._forward_box(features, proposals) # 框分支 mask_features = [features[f] for f in model.roi_heads.in_features] mask_features = model.roi_heads.mask_pooler(mask_features, [x.pred_boxes for x in instances])

需要说明的是,ImageList.from_tensors前的预处理、model.backbone的输入张量等细节都必须以你所用模型的实际forward()代码为准(例如 R-CNN 系列在GeneralizedRCNN元架构内还会先做preprocess_image归一化)。文档明确提醒:无论选哪条路径,都要先读懂现有 forward 代码,才能写出正确的取数逻辑。

七、在 OOTDiffusion 中的落地:人体解析模块的模型使用范例

上述机制在 OOTDiffusion 的 preprocess/humanparsing 模块中有非常直接的工程化应用——该模块用 detectron2 在 CIHP 数据集上微调并推理人体解析 Mask R-CNN 模型。

配置即架构:两份关键 YAML

parsing_finetune_cihp.yaml 是微调配置,其关键项正是前面各节讨论的cfg字段:

_BASE_: "cascade_mask_rcnn_X_152_32x8d_FPN_IN5k_gn_dconv.yaml" MODEL: MASK_ON: True WEIGHTS: "model_0039999_e76410.pkl" # 初始权重(.pkl,Caffe2 兼容加载) ROI_HEADS: NUM_CLASSES: 1 SOLVER: IMS_PER_BATCH: 16 STEPS: (140000, 180000) MAX_ITER: 200000 BASE_LR: 0.02 INPUT: MIN_SIZE_TRAIN: (640, 864) MIN_SIZE_TRAIN_SAMPLING: "range" MAX_SIZE_TRAIN: 1440 CROP: {ENABLED: True} DATASETS: TRAIN: ("CIHP_train",) TEST: ("CIHP_val",) OUTPUT_DIR: "./finetune_output"

而 parsing_inference.yaml 是推理配置,注意它把微调产物./finetune_ouput/model_final.pth作为MODEL.WEIGHTS,并额外设置了 NMS 与得分阈值:

_BASE_: "cascade_mask_rcnn_X_152_32x8d_FPN_IN5k_gn_dconv.yaml" MODEL: MASK_ON: True WEIGHTS: "./finetune_ouput/model_final.pth" ROI_HEADS: NMS_THRESH_TEST: 0.95 SCORE_THRESH_TEST: 0.5 NUM_CLASSES: 1 SOLVER: IMS_PER_BATCH: 1 STEPS: (30000, 45000) MAX_ITER: 50000 BASE_LR: 0.02 INPUT: MIN_SIZE_TRAIN: (640, 864) MIN_SIZE_TRAIN_SAMPLING: "range" MAX_SIZE_TRAIN: 1440 CROP: {ENABLED: True} TEST: AUG: {ENABLED: True} DATASETS: TRAIN: ("CIHP_trainval",) TEST: ("CIHP_test",) OUTPUT_DIR: "./inference_output"

这两份文件完整映射了本文的模型协议:NUM_CLASSES: 1意味着输出pred_classes的取值范围是[0, 1)(即二分类:人体/背景);MASK_ON: True保证输出 dict 含pred_masksSCORE_THRESH_TEST: 0.5NMS_THRESH_TEST: 0.95则对应demo.py--confidence-threshold注入的cfg.MODEL.ROI_HEADS.SCORE_THRESH_TEST

从 Mask R-CNN 输出到解析图:parsing_api.py的推理管线

真正面向 OOTDiffusion 实际使用的是 run_parsing.py 中的Parsing类——它加载checkpoints/humanparsing/parsing_atr.onnxparsing_lip.onnx两个 ONNX 模型(后者来自本文所述的 detectron2 解析模型导出),通过 onnxruntime 以 CPU/GPU 顺序执行模式推理。其核心逻辑 parsing_api.py 中的onnx_inference(L121-L185)展示了"模型输出 → 解析图"的完整后处理:

  1. 将模型输出的 logits 上采样到[512, 512],用transform_logits映射回原图坐标;
  2. np.argmax(logits_result, axis=2)取每个像素的类别——这正是本文第四节中sem_seg预测的逐像素 argmax 语义;
  3. 针对上衣(类别 4)与手臂(类别 14/15)做hole_fill空洞填充、refine_hole孔洞精修,得到干净的服装掩码;
  4. 融合 LIP 模型的颈部解析结果(neck_mask判定逻辑),用get_palette(19)生成 19 类调色板并输出带调色板的解析图(PIL.Image.putpalette)。

也就是说,detectron2 的模型输出协议(pred_masks/logits 张量)在 OOTDiffusion 中被进一步加工成了虚拟试穿所需的"服装区域掩码 + 人脸掩码(face_mask)",后者直接喂给 ootd 的扩散模型作为控制信号。

八、总结与延伸阅读

本文以 detectron2 官方模型使用文档为骨架,梳理了四条核心能力:构建build_model+ 注册表机制)、权重管理DetectionCheckpointer.pth/.pkl的差异化处理)、调用协议(训练态EventStorage与推理态DefaultPredictor)、数据协议list[dict]输入与instances/sem_seg/panoptic_seg输出),并给出了获取中间张量的两种方法。这些能力在 OOTDiffusion 的人体解析模块中均有真实落地的配置与代码可对照验证。

想继续深挖的读者可以沿着以下路径展开:

  • 注册表与自定义模型:write-models.md(对应 API 文档 modeling.rst);
  • 输入数据管线:data.rst、dataset_mapper.py;
  • 检查点 API:checkpoint.rst、detection_checkpoint.py;
  • 推理引擎:engine.rst、defaults.py;
  • 结构对象(Boxes/Instances/Masks/Keypoints):structures.rst;
  • 端到端示例:demo/demo.py 与 demo/predictor.py。

掌握这套模型使用范式后,无论是迁移到新数据集微调解析模型,还是替换骨架网络、自定义输出分支,都能在 detectron2 的配置与注册体系内以最小代价完成。

【免费下载链接】OOTDiffusion[AAAI 2025] Official implementation of "OOTDiffusion: Outfitting Fusion based Latent Diffusion for Controllable Virtual Try-on"项目地址: https://gitcode.com/GitHub_Trending/oo/OOTDiffusion

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

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

STM32F103 CAN1重映射原理与引脚选型实战指南

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

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

Uncloud 集群系统服务日志排查指南:uc machine logs 详解

Uncloud 集群系统服务日志排查指南:uc machine logs 详解 【免费下载链接】uncloud A lightweight tool for deploying and managing containerised applications across a network of Docker hosts. Bridging the gap between Docker and Kubernetes ✨ 项目地址…

作者头像 李华
网站建设 2026/9/17 23:46:45

8款AI论文写作工具评测与使用指南

1. AI论文写作工具的价值与现状作为一名在学术圈摸爬滚打多年的研究者,我深刻理解论文写作过程中的痛点。从文献综述到数据呈现,从格式排版到查重降重,每个环节都耗费大量时间精力。而AI写作工具的出现,正在改变这一现状。目前市面…

作者头像 李华