☰
IDM-VTON 中的 Detectron2 模型部署导出指南:TorchScript / ONNX / Caffe2 转换全解析
2026/10/3 8:29:48 网站建设 项目流程
  • 计算机视觉
  • 深度学习
  • 媒体生成

【免费下载链接】IDM-VTON

[ECCV2024] IDM-VTON : Improving Diffusion Models for Authentic Virtual Try-on in the Wild

项目地址:https://gitcode.com/GitHub_Trending/id/IDM-VTON
点击查看免费下载

导读:本指南围绕 gradio_demo/detectron2/export/README.md 展开,系统讲解 IDM-VTON 仓库内嵌的 detectron2 部署工具链如何将训练好的检测/分割模型转换为 TorchScript、ONNX 与(已弃用的)Caffe2 三种可部署格式。读完你将掌握导出 API 的完整调用链(Caffe2Tracer→Caffe2Model)、Caffe2 兼容模型的底层改写机制、TorchScript 脚本化Instances对象的特殊处理,以及模型校验、持久化与可视化排错的具体方法——这套能力对任何需要把 PyTorch 视觉模型送入生产推理引擎(TorchScript Runtime、ONNX Runtime、Caffe2)的工程实践都直接可复用。

1. 该模块在整个项目中的位置

IDM-VTON 是一个虚拟试穿(Virtual Try-on)项目,其gradio_demo目录内置了一个完整的 detectron2 及其 DensePose 扩展(gradio_demo/densepose),用于人体解析、DensePose 姿态估计等前置任务。export子目录(gradio_demo/detectron2/export)正是 detectron2 官方“模型部署导出”代码的嵌入式副本,它的职责非常单一且明确:把训练/推理用的 PyTorch 模型,改写成可被工业推理引擎直接加载的序列化模型。

原文档给出的定位如下:

  • 该目录包含一套将 detectron2 模型“准备用于部署”的代码;
  • 目前支持导出为TorchScript、ONNX与(已弃用)Caffe2三种格式;
  • 其使用方式指向外部部署文档(本文以仓库内源码为准,不再依赖外部链接)。

从源码结构看,gradio_demo/detectron2/export 由十个左右的模块组成,它们共同构成“重写模型 → 追踪/脚本化 → 导出 → 校验/可视化”的完整流水线:

模块职责
api.py对外的统一入口:Caffe2Tracer与Caffe2Model
caffe2_modeling.py把 detectron2 meta-arch 改写成“可被 Caffe2 追踪”的版本
caffe2_export.py导出 ONNX、Caffe2 protobuf,以及图可视化
shared.py共享工具:设备推断、图优化、protobuf 参数读写
caffe2_inference.py用 Caffe2 运行时模拟 detectron2 推理接口
c10.pyRPN / ROIPooler / 检测头等组件的 Caffe2 算子实现
torchscript.pyTorchScript 脚本化辅助与 IR 导出
torchscript_patch.py对不可脚本化类(Instances等)的动态补丁
flatten.py富结构输入/输出的张量化(Schema 机制)
caffe2_patch.py递归替换模型中特定组件为 Caffe2 兼容类

2. 三种导出格式的取舍:TorchScript / ONNX / Caffe2

原文档声明了三种目标格式,源码则进一步揭示了各自的实现深度与适用场景。

2.1 TorchScript:最贴近 PyTorch 生态的导出路径

TorchScript 导出通过Caffe2Tracer.export_torchscript()触发(api.py),其核心只有一步:

with torch.no_grad(): return torch.jit.trace(self.traceable_model, (self.traceable_inputs,))

即对“Caffe2 兼容版模型”做一次trace,得到torch.jit.TracedModule,可进一步调用.save()落盘。注意这里使用的是trace而非script,因为检测模型的输出中包含Instances等动态结构,直接脚本化困难。

当模型以script方式导出时,torchscript.py 提供了专门的辅助函数scripting_with_instances(model, fields):它会在脚本化前创建一个属性全部“静态化”的new_Instances类,并强制编译期编译器在遇到Instances时使用它,脚本化完成后自动还原进程状态。调用示例:

fields = {"proposal_boxes": Boxes, "objectness_logits": torch.Tensor} torchscript_model = scripting_with_instances(model, fields)

注意其前置约束——只支持evaluation mode下的模型(源码中以assert not model.training强制校验)。

scripting_with_instances的底层依赖 torchscript_patch.py 中的动态补丁:

  • patch_instances(fields):在临时目录中生成一个静态字段的脚本化Instances子类模块,通过_clear_jit_cache()清空 JIT 编译缓存后注册到 TorchScript 编译器,并支持__len__、to、__getitem__、cat、get_fields等常用方法;
  • freeze_training_mode(model):把各子模块的training属性标注为torch.jit.Final[bool]常量,让训练分支代码在脚本编译时被“元编译”裁剪掉;
  • patch_nonscriptable_classes():为ResNet、FPN注入__prepare_scriptable__,将其内部nn.Sequential替换为nn.ModuleList(规避 PyTorch 已知的脚本化缺陷),并把StandardROIHeads的mask_on/keypoint_on标注为常量。

此外,dump_torchscript_IR(model, dir)会把 TracedModule / ScriptModule 的代码与 IR 输出到目录,用于调试导出后的图结构,产出model_ts_code.txt、model_ts_IR.txt、model_ts_IR_inlined.txt、model.txt四个排错文件。

2.2 ONNX:经 Caffe2 兼容模型导出

ONNX 导出路径在 caffe2_export.py 的export_onnx_model中实现:

torch.onnx.export( model, inputs, f, operator_export_type=OperatorExportTypes.ONNX_ATEN_FALLBACK, )

几个值得注意的实现事实:

  1. 导出前会递归断言所有模块都处于eval模式(model.apply(_check_eval)),避免训练/推理状态不一致导致 ONNX 图错误;
  2. 采用ONNX_ATEN_FALLBACK算子导出类型,允许某些无法标准化的 PyTorch 算子以ATen形式落入 ONNX 图中;
  3. api.py 的export_onnx()文档明确警告:经此路径导出的 ONNX 模型含 Caffe2 专属自定义算子,无法被 onnxruntime 或 TensorRT 直接执行,如需对接这些运行时,需要额外的后处理/变换 pass,而本项目并不提供。

因此,从源码可以推断:当前仓库中 ONNX 是通往 Caffe2 的中间表示(见下文 2.3),而非面向通用 ONNX Runtime 的最终产物。

2.3 Caffe2:经 ONNX 中转的完整导出链路(已弃用)

Caffe2 导出的完整链路如下(对应 caffe2_export.py 的export_caffe2_detection_model):

  1. 深拷贝模型并断言其具备encode_additional_info方法(Caffe2 兼容 meta-arch 的标志);
  2. 先走 ONNX 导出(export_onnx_model),日志明确提示“ONNX 的一些警告是预期的,通常无需担心”;
  3. 用Caffe2Backend.onnx_graph_to_caffe2_net(onnx_model)把 ONNX 图转为 Caffe2 的init_net与predict_net两个 protobuf;
  4. 依次执行图优化 pass:
    • fuse_alias_placeholder:移除追踪期插入的AliasWithName占位算子;
    • GPU 输入时执行fuse_copy_between_cpu_and_gpu(合并多余的 CPU/GPU 拷贝算子)、remove_dead_end_ops、_assign_device_option(为每个算子标注设备选项);
    • remove_reshape_for_fc:Caffe2 的 FC 算子原生支持 4D 张量,因此可移除 ONNX 导出时为nn.Linear插入的动态 Reshape 子图;
    • group_norm_replace_aten_with_caffe2:把 ONNX 图中的ATen group_norm原位替换为 Caffe2 的GroupNorm算子;
  5. 调用model.encode_additional_info(predict_net, init_net)把推理所需的元信息(如size_divisibility、device、meta_architecture)写入 protobuf 参数;
  6. 输出每个网络的算子统计表,便于人工审查导出质量。

export_caffe2_detection_model返回(predict_net, init_net),交给Caffe2Model包装后即可在 PyTorch 侧以“假 nn.Module”的形式驱动(见第 4 节)。

3. 核心导出入口:Caffe2Tracer 与 Caffe2Model

api.py 定义了两个对外核心类,也是原文档所述“部署准备”能力的最直接体现。

3.1 Caffe2Tracer:构造可追踪的 Caffe2 兼容模型

Caffe2Tracer.__init__做三件事:

  1. 校验cfg必须是CfgNode、模型必须是torch.nn.Module;
  2. 根据cfg.MODEL.META_ARCHITECTURE从META_ARCH_CAFFE2_EXPORT_TYPE_MAP查表,用深拷贝的原始模型构造对应的 Caffe2 兼容 meta-arch;
  3. 通过get_caffe2_inputs(inputs)把 PyTorch 风格的输入转换为 Caffe2 风格的两个张量。

其导出后的计算图固定接收两个输入张量(docstring 明确给出):

  • data:(1, C, H, W)的 float 图像,取值通常在[0, 255];(H, W)通常需按模型结构补齐到 32 的倍数;
  • im_info:N×3的 float 张量,每行为(height, width, 1.0),其中 height/width 是补齐前的真实图像尺寸。

同时,Caffe2Tracer只支持内建的 meta-arch(源码META_ARCH_CAFFE2_EXPORT_TYPE_MAP目前只有GeneralizedRCNN与RetinaNet两项),且不支持 batch 推理(代码注释中欢迎社区贡献)。

3.2 Caffe2Model:protobuf 模型的 PyTorch 风格包装

Caffe2Model是一个包裹 Caffe2 protobuf 的nn.Module,始终处于eval模式,提供以下方法:

  • save_protobuf(output_dir):落盘三个文件——model.pb(图定义)、model_init.pb(模型参数)、model.pbtxt(人类可读的图定义,部署时不需要);
  • load_protobuf(dir):从model.pb+model_init.pb反向加载;
  • save_graph(output_file, inputs=None):把网络导出为 SVG 图;若提供输入,则会实际运行网络并记录每个 blob 的 shape,把 shape 信息一起画进图里,用于可视化排错;
  • __call__(inputs):通过ProtobufDetectionModel模拟 detectron2 模型的输入/输出格式,方便与原始 Torch 模型逐项对比结果(caffe2_inference.py)。

代码注释还提示:__call__因为包含 PyTorch/Caffe2 间的额外转换,不适合用来做性能基准测试——它服务于正确性校验而非性能测量。

4. Caffe2 兼容模型的底层改写机制

Caffe2 导出能成立,前提是先把原始模型“翻译”成只含 Caffe2 算子的版本。这条链路由三个文件协同完成。

4.1 Caffe2 兼容 meta-arch 基类

caffe2_modeling.py 定义了Caffe2MetaArch基类,其关键设计是forward必须“可追踪、且只用 Caffe2 兼容算子”,这样被 trace 的图才能转换成 Caffe2 图。它还把输入预处理(归一化、按size_divisibility批量补齐)收敛到统一的_caffe2_preprocess_image,并将pixel_mean/pixel_std归一化直接编入图中(图输入即为未归一化的data)。

典型实现是Caffe2GeneralizedRCNN,它的 trace forward 依次执行:预处理 → backbone 提特征 → proposal 生成 → ROI heads 推理 →tuple(detector_results[0].flatten())扁平化输出。与此同时:

  • patch_generalized_rcnn(caffe2_patch.py)递归地把模型里的rpn.RPN换成Caffe2RPN、poolers.ROIPooler换成Caffe2ROIPooler;
  • ROIHeadsPatcher会把 Fast R-CNN / Mask R-CNN / Keypoint R-CNN 的推理函数 mock 成 Caffe2 版(分别对应Caffe2FastRCNNOutputsInference、Caffe2MaskRCNNInference、Caffe2KeypointRCNNInference),并受配置项EXPORT_CAFFE2.USE_HEATMAP_MAX_KEYPOINT控制是否输出稀疏关键点。
  • encode_additional_info会把size_divisibility、device、meta_architecture等元信息写进 protobuf,供后续推理时读取。

4.2 算子级 Caffe2 实现

c10.py 提供了检测流水线关键环节的 Caffe2 算子实现,例如:

  • Caffe2RPN._generate_proposals使用torch.ops._caffe2.GenerateProposals/CollectRpnProposals完成多尺度 RPN 提案合并,并在pre_nms_topk/post_nms_topk/nms_thresh/min_size等参数上与 detectron2 对齐;
  • Caffe2ROIPooler.forward在 FPN 多级场景下用DistributeFpnProposals+ 各级RoIAlign+BatchPermutation还原原始顺序;
  • caffe2_fast_rcnn_outputs_inference用BBoxTransform+BoxWithNMSLimit完成检测头推理与 NMS 后处理,其 softmax/sigmoid 分类分支、旋转框(RotatedBoxes)分支都做了兼容处理。

由此可以推断:Caffe2 导出并不等于“全图展开成简单算子”,而是尽量把高层的 RPN、NMS、RoIAlign 等逻辑映射到 Caffe2 原生复合算子,从而保证执行效率与语义一致。

4.3 共享图工具:protobuf 读写与设备推断

shared.py 是整个 export 模块的“工具箱”,涵盖:

  • protobuf 参数读写:get_pb_arg_valf/get_pb_arg_vali/get_pb_arg_vals/get_pb_arg_floats及写入用的check_set_pb_arg,支持 float/int/string/floats 等参数类型;
  • 设备类型静态推断:infer_device_type基于 SSA 形式前后向传播已知 blobs 的 CPU/GPU 状态,并通过CopyCPUToGPU/CopyGPUToCPU的更新规则处理跨设备拷贝;
  • 算子/参数图变换:fuse_alias_placeholder、rename_op_input/output、remove_reshape_for_fc、fuse_copy_between_cpu_and_gpu等;
  • 图导出与可视化:save_graph通过 pydot 生成 PNG/PDF/SVG,可带 blob shape 标注;
  • workspace 管理:ScopedWS上下文管理器隔离 Caffe2 工作空间,避免不同模型实例相互污染。

5. 富结构输入/输出如何被“张量化”:Schema 与 TracingAdapter

detectron2 的模型输入是list[dict](每张图一个 dict,含image等字段),输出是Instances等富结构对象,而torch.jit.trace只接受/产生张量元组。这一矛盾由 flatten.py 解决:

  • Schema.flatten(obj)把任意对象压平成张量元组,同时记录一张“重建蓝图”(可序列化的 dataclass),随后用schema(flattened_values)即可还原原对象;
  • 支持的 schema 类型覆盖IdentitySchema、ListSchema、TupleSchema、DictSchema、InstancesSchema、TensorWrapSchema(用于Boxes、RotatedBoxes、ROIMasks等张量包装类);
  • TracingAdapter是连接“富接口模型”与“trace 接口”的适配器:它负责把输入压平、在no_grad与patch_builtin_len()(把内置len替换为__len__,避免追踪期生成 ONNX 常量)下执行推理,再压平输出并缓存outputs_schema。有了它,任何 detectron2 模型都能被torch.jit.trace,并能通过 schema 把扁平输出还原为原始结构用于比对。

TracingAdapter还支持allow_non_tensor=True的宽松模式:此时仅保留张量用于追踪、丢弃 schema 重建能力,适合只关心单次执行图(例如 FLOP 统计)的场景。

6. 校验、持久化与可视化排错

6.1 导出结果与原模型的正确性对比

Caffe2Model.__call__的设计初衷之一就是与原始 Torch 模型逐项比对输出。文档注释给出了最小用法:

c2_model = Caffe2Tracer(cfg, torch_model, inputs).export_caffe2() inputs = [{"image": img_tensor_CHW}] outputs = c2_model(inputs) orig_outputs = torch_model(inputs)

其背后由 caffe2_inference.py 的ProtobufDetectionModel支撑:它包装ProtobufModel运行 Caffe2 网络,从 protobuf 参数中读取size_divisibility、device、meta_architecture,再通过对应 meta-arch 的get_outputs_converter把 Caffe2 的扁平输出还原成 detectron2 标准格式(list[Instances])。

6.2 protobuf 落盘与加载

部署时只需三个文件的model.pb+model_init.pb(model.pbtxt供人工阅读)。加载方Caffe2Model.load_protobuf(dir)要求目录中包含这两个二进制文件,并从它们反序列化出网络定义与参数。

6.3 图可视化

save_graph支持 SVG/PNG/PDF 输出,run_and_save_graph(caffe2_export.py)则会在给定输入上实际运行init_net+predict_net,收集每个 numpy blob 的 shape,连同算子依赖关系一起渲染到图中,是排查“算子缺失 / 尺寸不匹配”类导出问题的高效手段。若网络运行中途抛错,该函数会捕获RuntimeError并以 warning 记录,仍继续保存已有的 blob shape 信息。

7. 在 IDM-VTON 中实践导出:适用前提与限制

结合仓库实际情况,本模块的使用存在以下必须知晓的前提与限制:

  1. 支持范围:META_ARCH_CAFFE2_EXPORT_TYPE_MAP仅覆盖GeneralizedRCNN与RetinaNet两种内建 meta-arch(见 caffe2_modeling.py),自定义结构需自行扩展;导出要求追踪输入能产生有效检测结果(源码断言“没有检测结果的追踪会导致错误 trace”)。
  2. Caffe2 已弃用:原文档明确标注 Caffe2 为 deprecated;ONNX 导出产物含 Caffe2 专属算子,不能直接被 onnxruntime / TensorRT 运行。实际可无缝落地的是 TorchScript 路径(Caffe2Tracer.export_torchscript()+.save())或 Caffe2 生态自身。
  3. 单图推理:Caffe2Tracer不支持 batch 推理,追踪图以(data, im_info)两个固定张量为输入,图像需按 backbone 的size_divisibility补齐(常见为 32 的倍数)。
  4. 部分算子仅有 CPU 实现:源码 docstring 指出 Caffe2 中部分算子没有 GPU 实现,跨设备部署前需核对该问题。

读者可继续阅读仓库中的以下文件深入研习:统一入口 api.py、导出流水线 caffe2_export.py、兼容模型构造 caffe2_modeling.py、算子级实现 c10.py、共享工具 shared.py,以及脚本化辅助 torchscript.py 与 torchscript_patch.py。

  • 计算机视觉
  • 深度学习
  • 媒体生成

【免费下载链接】IDM-VTON

[ECCV2024] IDM-VTON : Improving Diffusion Models for Authentic Virtual Try-on in the Wild

项目地址:https://gitcode.com/GitHub_Trending/id/IDM-VTON
点击查看免费下载

相关推荐

上一篇:Ninja从入门到上手:5分钟安装并编译你的第一个项目,快速开始完整教程
下一篇:seaborn 声明式绘图接口 Plot 全解:从 autosummary 模板到源码级方法体系与 Plot.config 配置机制

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

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询