- 计算机视觉
- 深度学习
- 媒体生成
【免费下载链接】IDM-VTON
[ECCV2024] IDM-VTON : Improving Diffusion Models for Authentic Virtual Try-on in the Wild
导读:本文围绕 IDM-VTON 仓库内置的 detectron2 命令行工具集(
preprocess/humanparsing/mhp_extension/detectron2/tools/)展开,系统讲解train_net.py、plain_train_net.py、benchmark.py、visualize_json_results.py、visualize_data.py五个核心脚本的定位、用法与底层实现,并结合本仓库人体解析(Human Parsing)预处理链路中实际使用的配置文件与 shell 脚本,说明如何用这套工具完成模型训练、推理评测、速度基准测试与数据可视化,帮助你在虚拟试穿(Virtual Try-on)项目中快速上手 detectron2 的命令行工作流。
一、tools 目录在 IDM-VTON 中的作用
IDM-VTON 的虚拟试穿流水线在预处理阶段需要精确的人体解析(human parsing)结果——即把人体图像分割为衣服、裤子、头发、皮肤等语义区域,作为后续 Diffusion 模型的条件输入。该能力由preprocess/humanparsing/下的 mhp_extension 提供,其中内嵌了一份完整的 detectron2 代码树,而preprocess/humanparsing/mhp_extension/detectron2/tools/正是面向命令行的入口集合。
从仓库目录结构看,tools/ 下实际包含以下可执行脚本:
| 脚本 | 用途 |
|---|---|
| train_net.py | 基于DefaultTrainer的标准训练 / 评测入口 |
| plain_train_net.py | 手写训练循环的"极简版"训练脚本 |
| benchmark.py | 训练 / 推理 / 数据加载速度基准测试 |
| visualize_json_results.py | 将 COCO/LVIS 格式 JSON 评测结果可视化为图片 |
| visualize_data.py | 可视化标注原始数据或经预处理 / 增强后的训练数据 |
| analyze_model.py | 统计模型的 FLOPs、激活量、参数量与结构 |
| finetune_net.py | 微调网络(本仓库人体解析训练 / 推理实际使用) |
| convert-torchvision-to-d2.py | 将 torchvision 预训练权重转换为 detectron2 格式 |
| inference.sh 与 run.sh | 解析模型的推理与微调一键脚本 |
| deploy/ | Caffe2 / TorchScript 部署转换相关脚本与说明 |
其中train_net.py、plain_train_net.py、benchmark.py、visualize_json_results.py、visualize_data.py是官方 README 明确讲解的五个通用工具,下面逐一展开。
二、train_net.py:面向内置模型的标准训练入口
2.1 脚本定位
train_net.py是一个"读取配置并执行训练或评测"的通用入口脚本,官方在 README 中明确说明:它是为训练 detectron2 内置(builtin)模型而设计的,因此脚本内包含许多与这些内置模型绑定的逻辑(如按数据集元数据自动选择评估器),如果你的项目需求比较特殊,官方建议把 detectron2 当作库来使用,并以本脚本为 API 调用示例。
2.2 工作流:setup → Trainer → launch
从源码看,脚本的执行链路非常清晰(train_net.py):
setup(args):通过get_cfg()创建默认配置 →cfg.merge_from_file(args.config_file)合并 YAML 配置 →cfg.merge_from_list(args.opts)合并命令行覆盖项 →cfg.freeze()冻结 →default_setup(cfg, args)完成日志、随机种子等基础设置;main(args)分支处理:- 评测模式(
--eval-only):Trainer.build_model(cfg)构建模型,DetectionCheckpointer(...).resume_or_load(cfg.MODEL.WEIGHTS)加载权重,Trainer.test评测;若cfg.TEST.AUG.ENABLED为真,还会追加 TTA(测试时增强)评测并输出_TTA后缀的结果; - 训练模式:实例化
Trainer(cfg),resume_or_load支持断点续训,trainer.train()启动训练;同样支持TEST.AUG.ENABLED时注册EvalHook周期性做 TTA 评测;
- 评测模式(
- 入口统一走
launch(...)多机多卡启动函数,传入args.num_gpus、num_machines、machine_rank、dist_url。
2.3 Trainer 子类:按数据集自动选择评估器
脚本中的Trainer(DefaultTrainer)只重写了build_evaluator和test_with_TTA两个类方法(train_net.py)。build_evaluator的核心逻辑是读取MetadataCatalog.get(dataset_name).evaluator_type,据此分发:
sem_seg/coco_panoptic_seg→SemSegEvaluator;coco/coco_panoptic_seg→COCOEvaluator(panoptic 还叠加COCOPanopticEvaluator);cityscapes_instance/cityscapes_sem_seg→ Cityscapes 系列评估器(要求 GPU 数量不小于 rank);pascal_voc→PascalVOCDetectionEvaluator;lvis→LVISEvaluator。
评测输出默认写入cfg.OUTPUT_DIR/inference;TTA 评测写入inference_TTA子目录。
2.4 命令行用法(官方 README + GETTING_STARTED 原文继承)
在配置好数据集(见 datasets/README.md)后,8 卡训练:
cd tools/ ./train_net.py --num-gpus 8 \ --config-file ../configs/COCO-InstanceSegmentation/mask_rcnn_R_50_FPN_1x.yaml官方配置默认按 8 卡设计。改为 1 卡训练时需同步调整学习率与 batch size(detectron2 官方建议线性缩放学习率):
./train_net.py \ --config-file ../configs/COCO-InstanceSegmentation/mask_rcnn_R_50_FPN_1x.yaml \ --num-gpus 1 SOLVER.IMS_PER_BATCH 2 SOLVER.BASE_LR 0.0025注意:对大多数模型,detectron2不支持 CPU 训练。
仅评测(--eval-only配合权重路径):
./train_net.py \ --config-file ../configs/COCO-InstanceSegmentation/mask_rcnn_R_50_FPN_1x.yaml \ --eval-only MODEL.WEIGHTS /path/to/checkpoint_file更多选项查看./train_net.py -h。
三、plain_train_net.py:手写训练循环的极简替代
3.1 与 train_net.py 的区别
plain_train_net.py与train_net.py功能等价(都能训练标准模型),但不再使用Trainer,而是把训练循环完整展开。README 的评价是"功能更少,但对想深度修改逻辑的开发者更友好(more friendly to hackers)"。
对比源码可发现具体差异:
- 训练循环显式化:
do_train中依次手动执行model.train()→build_optimizer→build_lr_scheduler→DetectionCheckpointer(带 optimizer 与 scheduler,支持断点续训)→for data, iteration in zip(data_loader, range(start_iter, max_iter))逐轮迭代,手动losses.backward()/optimizer.step()/scheduler.step()(plain_train_net.py); - 评测独立成函数:
do_test对cfg.DATASETS.TEST中每个数据集构建测试 loader 与评估器,inference_on_dataset产出结果并以 CSV 格式打印(plain_train_net.py); - 功能取舍:源码注释明确说明"不支持精确计时与精确 BN(precise BN)",训练循环内
loss_dict经comm.reduce_dict聚合后写入CommonMetricPrinter/JSONWriter(输出metrics.json)/TensorboardXWriter三类事件写入器。
3.2 适用场景
如果你需要给人体解析模型插入自定义的数据增强、损失函数或调度逻辑,从plain_train_net.py复制框架再改动远比逆向理解DefaultTrainer的 hook 机制更省事。这正是官方在脚本 docstring 中推荐的做法:把 detectron2 当作库,自定义自己的训练循环。
四、benchmark.py:训练 / 推理 / 数据加载三合一测速
4.1 基本用法(README 原文继承)
python benchmark.py --config-file config.yaml --task train/eval/data [optional DDP flags]--task为必选参数(choices 限定train/eval/data三选一,见 benchmark.py):
data:只测数据加载速度;train:测训练速度(源码注释特别提醒:训练速度可能不具备代表性,R-CNN 模型的训练开销随数据内容与模型质量变化);eval:只测单卡推理,脚本内部断言num_gpus == 1 and num_machines == 1。
4.2 源码级细节
- 额外依赖:脚本头部注明需要
psutil(用于打印 RAM 占用:psutil.virtual_memory()输出已用/总内存,单位 GB); - data 任务:
build_detection_train_loader(cfg)构建 loader,先 warmup 10 轮记录启动时间(startup time),再跑 1000 轮计时,并额外做 10 轮重复计时取稳定性参考; - train / eval 任务:源码将
cfg.DATALOADER.NUM_WORKERS置为 0,并先把 loader 取出的前 100 个 batch 固化为dummy_data(DatasetFromList)循环供给,从而在固定输入下衡量模型吞吐,避免 IO 抖动干扰;eval 前还做 5 轮 warmup。setup 阶段还会强制cfg.SOLVER.BASE_LR = 0.001(避免 NaN,注释说明该值在本脚本中无实际意义)。
这一设计非常实用:调优人体解析模型的 batch size / dataloader worker 数量时,benchmark.py --task data可以快速定位数据管线的瓶颈。
五、visualize_json_results.py:评测结果的 JSON 可视化
5.1 基本用法(README 原文继承)
python visualize_json_results.py --input x.json --output dir/ --dataset coco_2017_val参数说明:--input指定由COCOEvaluator或LVISEvaluator落盘的 JSON 结果文件(必填);--output为输出目录(必填);--dataset指定数据集名,默认coco_2017_val;另有--conf-threshold控制置信度过滤阈值,默认0.5(visualize_json_results.py)。
若你的数据集不是 detectron2 内置数据集,README 明确提示需要自写注册脚本或直接修改本脚本(主要改动点就是
DatasetCatalog.get(args.dataset)与MetadataCatalog.get(args.dataset)所依赖的数据集注册)。
5.2 工作原理解读
- 读取 JSON 后按
image_id聚合预测(pred_by_image = defaultdict(list)); - 通过
DatasetCatalog.get(args.dataset)拿到每张图的原始标注 dict; - 类别 ID 映射逻辑:优先用元数据
thing_dataset_id_to_contiguous_id;否则若数据集名含lvis,按ds_id - 1映射(LVIS 类别 ID 约定);否则报错; create_instances把预测 dict 转成Instances结构:按置信度阈值筛选 → bbox 由BoxMode.XYWH_ABS转BoxMode.XYXY_ABS→ 重建Boxes/pred_classes/ 可选的pred_masks;- 用
Visualizer分别绘制预测框(draw_instance_predictions)与真实标注(draw_dataset_dict),左右拼接成一张图(np.concatenate(..., axis=1))后写入输出目录,方便直接对比模型输出与 ground truth。
六、visualize_data.py:标注与预处理数据的可视化
6.1 基本用法(README 原文继承)
python visualize_data.py --config-file config.yaml --source annotation/dataloader --output-dir dir/ [--show]--source(必填,二选一):annotation:直接可视化数据集的原始标注;dataloader:可视化经预处理 / 增强之后的训练数据(即真正喂给模型的 batch);
--config-file:配置文件路径(可选,若不填则只有命令行 opts);--output-dir:输出目录,默认./;--show:是否在 OpenCV 窗口中即时显示(启用时可视化缩放倍率scale = 2.0,否则为1.0)。
6.2 两个模式的实现差异
- dataloader 模式:
build_detection_train_loader(cfg)构造训练 loader,逐 batch 取出per_image["image"](PyTorch 的 C,H,W 张量),permute(1,2,0)转回 H,W,C 后用utils.convert_image_to_rgb(img, cfg.INPUT.FORMAT)按配置的输入格式还原 RGB,再调用visualizer.overlay_instances叠加 gt 框 / 掩码 / 关键点(visualize_data.py); - annotation 模式:遍历
cfg.DATASETS.TRAIN下所有数据集的标注 dict,用visualizer.draw_dataset_dict(dic)绘制;若cfg.MODEL.KEYPOINT_ON开启,还会先经filter_images_with_few_keypoints(dicts, 1)过滤掉关键点过少的样本。
6.3 重要注意事项
README 特别警告:使用--source dataloader时脚本不会自行停止,因为训练 dataloader 通常是无限循环的。实践中应配合--show逐张查看,或搭配管道限流(如head截断)使用。
七、IDM-VTON 实战:这套工具在本仓库解析链路中的真实用法
本仓库并未停留在工具展示层面——preprocess/humanparsing/mhp_extension/detectron2/tools/下的两个 shell 脚本与configs/Misc/中的解析配置,构成了 IDM-VTON 人体解析模型的实际训练 / 推理入口。
7.1 解析模型推理:inference.sh
inference.sh 内容如下:
python finetune_net.py \ --num-gpus 1 \ --config-file ../configs/Misc/parsing_inference.yaml \ --eval-only MODEL.WEIGHTS ./model_final.pth TEST.AUG.ENABLED False它复用finetune_net.py(基于 train_net 系列的微调变体),单卡加载 parsing_inference.yaml 做纯评测。该配置的关键点:
- 继承
cascade_mask_rcnn_X_152_32x8d_FPN_IN5k_gn_dconv.yaml的骨干结构(Cascade Mask R-CNN + X-152 + FPN + GN + DCN),MASK_ON: True; ROI_HEADS.NUM_CLASSES: 1:人体解析任务被建模为单类实例分割(只区分人体区域);NMS_THRESH_TEST: 0.95、SCORE_THRESH_TEST: 0.5:推理阶段的 NMS 与置信度阈值;SOLVER.IMS_PER_BATCH: 1、MAX_ITER: 50000、BASE_LR: 0.02、STEPS: (30000, 45000):微调阶段的调度参数;- 数据集为
CIHP_trainval/CIHP_test;输出目录./inference_output。
命令行中TEST.AUG.ENABLED False显式关闭 TTA,保证推理速度与结果确定性。
7.2 解析模型微调:run.sh
run.sh 内容如下:
python finetune_net.py \ --config-file ../configs/Misc/parsing_finetune_cihp+vip.yaml \ --num-gpus 8对应 8 卡微调配置 parsing_finetune_cihp.yaml 中:IMS_PER_BATCH: 16、MAX_ITER: 200000、STEPS: (140000, 180000)、BASE_LR: 0.02,输入短边在(640, 864)范围内随机采样(MIN_SIZE_TRAIN_SAMPLING: "range"),最长边 1440,并开启随机裁剪(CROP.ENABLED: True),即通过 resize + crop 增强适配 CIHP 数据集。可见官方 tools 工作流被完整复用到了本项目的人体解析模型训练中。
7.3 与整条预处理链路的衔接
在本仓库中,解析结果由 run_parsing.py 及 parsing_api.py 包装,供虚拟试穿管线调用(如配合 utils_mask.py 生成衣物掩码)。tools/中的脚本正是这一能力训练与调优阶段的底层支撑——理解它们,就能完全掌控从数据可视化、模型微调到批量推理的每一步。
八、命令速查表与选型建议
| 需求 | 推荐脚本 | 关键参数 |
|---|---|---|
| 标准训练 / 评测 | train_net.py | --config-file、--num-gpus、--eval-only MODEL.WEIGHTS ... |
| 自定义训练循环 | plain_train_net.py | 同上(手动管理 optimizer / scheduler / checkpoint) |
| 速度基准测试 | benchmark.py | --task train\|eval\|data(需psutil) |
| 评测 JSON 可视化 | visualize_json_results.py | --input、--output、--dataset、--conf-threshold |
| 数据 / 标注可视化 | visualize_data.py | --source annotation\|dataloader、--show |
| 模型结构 / 开销分析 | analyze_model.py | --tasks flop,activation,parameter,structure、--num-inputs |
| 解析模型训练 / 推理 | finetune_net.py+inference.sh/run.sh | 搭配 parsing_inference.yaml 等配置 |
选型建议:若你只是复用 IDM-VTON 现有解析模型做推理,直接执行inference.sh即可;若要在 CIHP 等自有数据集上重训解析头,以run.sh为起点并修改 parsing_finetune_cihp.yaml 中的数据集与调度参数;若需排查增强是否合理,优先用visualize_data.py --source dataloader逐张确认预处理后的训练样本。
九、延伸阅读与相关源码
- 工具集官方用法说明:tools/README.md
- 完整上手指南(训练、评测、演示 Demo 参数):GETTING_STARTED.md
- 模型仓库说明:MODEL_ZOO.md
- 数据集注册说明:data/datasets/README.md
- 人体解析整体入口:run_parsing.py 与 parsing_api.py
- 部署转换工具:tools/deploy/README.md
- 计算机视觉
- 深度学习
- 媒体生成
【免费下载链接】IDM-VTON
[ECCV2024] IDM-VTON : Improving Diffusion Models for Authentic Virtual Try-on in the Wild
相关推荐
Astro 集成 Svelte 指南:@astrojs/svelte 的安装配置、Svelte 5 支持与版本演进全解析
Astro 集成 Svelte 指南:@astrojs/svelte 的安装配置、Svelte 5 支持与版本演进全解析 @astrojs/svelte 是 A
计算机视觉深度学习媒体生成IDM-VTON 人体解析评测指南:深入解读 detectron2 的 DatasetEvaluator 与 inference_on_dataset 评估机制
IDM VTON 人体解析评测指南:深入解读 detectron2 的 DatasetEvaluator 与 inference_on_dataset 评估机制
计算机视觉深度学习媒体生成IDM-VTON 人像解析栈中的 detectron2 基准评测:Mask R-CNN 训练吞吐量对比与复现指南
IDM VTON 人像解析栈中的 detectron2 基准评测:Mask R CNN 训练吞吐量对比与复现指南 本文面向在 IDM VTON 中从事人体解析(
计算机视觉深度学习媒体生成
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考