- 计算机视觉
- 深度学习
- 媒体生成
【免费下载链接】IDM-VTON
[ECCV2024] IDM-VTON : Improving Diffusion Models for Authentic Virtual Try-on in the Wild
导读:本指南围绕 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.py | RPN / ROIPooler / 检测头等组件的 Caffe2 算子实现 |
| torchscript.py | TorchScript 脚本化辅助与 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, )几个值得注意的实现事实:
- 导出前会递归断言所有模块都处于
eval模式(model.apply(_check_eval)),避免训练/推理状态不一致导致 ONNX 图错误; - 采用
ONNX_ATEN_FALLBACK算子导出类型,允许某些无法标准化的 PyTorch 算子以ATen形式落入 ONNX 图中; - 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):
- 深拷贝模型并断言其具备
encode_additional_info方法(Caffe2 兼容 meta-arch 的标志); - 先走 ONNX 导出(
export_onnx_model),日志明确提示“ONNX 的一些警告是预期的,通常无需担心”; - 用
Caffe2Backend.onnx_graph_to_caffe2_net(onnx_model)把 ONNX 图转为 Caffe2 的init_net与predict_net两个 protobuf; - 依次执行图优化 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算子;
- 调用
model.encode_additional_info(predict_net, init_net)把推理所需的元信息(如size_divisibility、device、meta_architecture)写入 protobuf 参数; - 输出每个网络的算子统计表,便于人工审查导出质量。
export_caffe2_detection_model返回(predict_net, init_net),交给Caffe2Model包装后即可在 PyTorch 侧以“假 nn.Module”的形式驱动(见第 4 节)。
3. 核心导出入口:Caffe2Tracer 与 Caffe2Model
api.py 定义了两个对外核心类,也是原文档所述“部署准备”能力的最直接体现。
3.1 Caffe2Tracer:构造可追踪的 Caffe2 兼容模型
Caffe2Tracer.__init__做三件事:
- 校验
cfg必须是CfgNode、模型必须是torch.nn.Module; - 根据
cfg.MODEL.META_ARCHITECTURE从META_ARCH_CAFFE2_EXPORT_TYPE_MAP查表,用深拷贝的原始模型构造对应的 Caffe2 兼容 meta-arch; - 通过
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 中实践导出:适用前提与限制
结合仓库实际情况,本模块的使用存在以下必须知晓的前提与限制:
- 支持范围:
META_ARCH_CAFFE2_EXPORT_TYPE_MAP仅覆盖GeneralizedRCNN与RetinaNet两种内建 meta-arch(见 caffe2_modeling.py),自定义结构需自行扩展;导出要求追踪输入能产生有效检测结果(源码断言“没有检测结果的追踪会导致错误 trace”)。 - Caffe2 已弃用:原文档明确标注 Caffe2 为 deprecated;ONNX 导出产物含 Caffe2 专属算子,不能直接被 onnxruntime / TensorRT 运行。实际可无缝落地的是 TorchScript 路径(
Caffe2Tracer.export_torchscript()+.save())或 Caffe2 生态自身。 - 单图推理:
Caffe2Tracer不支持 batch 推理,追踪图以(data, im_info)两个固定张量为输入,图像需按 backbone 的size_divisibility补齐(常见为 32 的倍数)。 - 部分算子仅有 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
相关推荐
Detectron2 模型部署导出指南:TorchScript / ONNX / Caffe2 的完整导出方案
Detectron2 模型部署导出指南:TorchScript / ONNX / Caffe2 的完整导出方案 本指南基于 Detectron2 仓库中的 de
人工智能计算机视觉深度学习机器学习IDM-VTON 人体解析预处理中的 detectron2 模型部署导出:基于 ONNX 的 Caffe2 格式转换实践指南
IDM VTON 人体解析预处理中的 detectron2 模型部署导出:基于 ONNX 的 Caffe2 格式转换实践指南 在 IDM VTON(ECCV 2
计算机视觉深度学习媒体生成Meteor accounts-facebook 包实战指南:在 Meteor 3 应用中集成 Facebook OAuth 登录
Meteor accounts facebook 包实战指南:在 Meteor 3 应用中集成 Facebook OAuth 登录 accounts faceb
人工智能大模型媒体生成本地部署深度学习
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考