- 计算机视觉
- 深度学习
- 媒体生成
【免费下载链接】IDM-VTON
[ECCV2024] IDM-VTON : Improving Diffusion Models for Authentic Virtual Try-on in the Wild
导读
本文以 IDM-VTON 仓库 TridentNet 项目说明 为骨架,完整讲解 TridentNet(Scale-Aware Trident Networks for Object Detection)在 Detectron2 框架下的实现、训练与评估流程。通过对照仓库中真实的配置文件与源码,你可以掌握并行多分支结构共享权重、不同感受野提取尺度感知特征的核心原理,理解 TridentNet-Fast 零额外参数提速推理的实现机制,并能在实际环境中复现 COCO 上的完整训练与评估命令。
TridentNet 核心思想:用统一表示能力生成尺度感知特征
TridentNet 是发表于 ICCV 2019 的经典检测工作(论文作者:Yanghao Li、Yuntao Chen、Naiyan Wang、Zhaoxiang Zhang,详见仓库 README 提供的 BibTeX 条目)。其核心目标是:生成具有"统一表示能力"(uniform representational power)的尺度特定特征图。
传统单分支检测网络在面对同一物体的不同尺度时,深层网络往往对某个尺度响应更好、对另一尺度则能力退化,导致小目标或大目标的检测精度不均衡。TridentNet 的解决方案是构造一个并行多分支架构:
- 每个分支共享完全相同的变换参数(即共享卷积权重);
- 每个分支拥有不同的感受野(通过不同空洞率 dilation 实现);
- 各分支分别关注不同尺度范围的目标,从而让"尺度"这个维度在特征提取阶段就被显式建模。
TridentNet-Fast:不增加参数与计算量的快速近似
仓库中实际落地的是TridentNet-Fast,它是 TridentNet 的快速近似版本。其思想非常巧妙:训练阶段让全部分支并行工作以获得尺度感知能力,推理阶段只使用其中一个分支,从而在不引入任何额外参数和计算成本的前提下,获得相比普通 Faster R-CNN 显著的精度提升。
这一策略的工程实现依赖一个关键配置项MODEL.TRIDENT.TEST_BRANCH_IDX:
- 设为
-1:推理时聚合所有分支的结果(通过 NMS 合并,精度最高但计算量最大); - 设为非负整数(仓库默认
1,即中间分支):推理时仅使用指定分支做快速推理(推荐使用中间分支,因为其感受野居中、对尺度覆盖最均衡)。
在 trident_conv.py 的forward中可以看到这一逻辑的直接体现:训练或test_branch_idx == -1时,对每个分支输入分别执行一次F.conv2d(各自使用自己的 padding 与 dilation);而 TridentNet-Fast 推理模式下,仅对第一个输入执行一次卷积,返回单元素列表:
if self.training or self.test_branch_idx == -1: outputs = [ F.conv2d(input, self.weight, self.bias, self.stride, padding, dilation, self.groups) for input, dilation, padding in zip(inputs, self.dilations, self.paddings) ] else: outputs = [ F.conv2d( inputs[0], self.weight, self.bias, self.stride, self.paddings[self.test_branch_idx], self.dilations[self.test_branch_idx], self.groups, ) ]注意其中self.weight是所有分支共享的单一权重张量(TridentConv只声明了一个nn.Parameter),这正是"共享变换参数、不同感受野"这一设计在代码层面的落点:kaiming_uniform_初始化权重,每个分支只是用不同的空洞率与 padding 对同一份权重做卷积。
源码架构解析:TridentNet 的五个核心组件
本仓库将 TridentNet 以 Detectron2 project 的形式完整实现,代码位于 projects/TridentNet,由五部分构成:
1. 配置注入:config.py
add_tridentnet_config(cfg)通过CfgNode为 Detectron2 注入MODEL.TRIDENT命名空间,所有 TridentNet 专属参数都在这里定义,下表给出参数名、默认值与语义:
| 配置项 | 默认值 | 语义 |
|---|---|---|
MODEL.TRIDENT.NUM_BRANCH | 3 | TridentNet 的分支数量 |
MODEL.TRIDENT.BRANCH_DILATIONS | [1, 2, 3] | 各分支对应的空洞率(dilation),与分支数一一对应 |
MODEL.TRIDENT.TRIDENT_STAGE | "res4" | 应用 Trident block 的 ResNet 阶段,按原论文默认取 Res4 |
MODEL.TRIDENT.TEST_BRANCH_IDX | 1 | TridentNet-Fast 推理分支索引;-1表示推理时聚合所有分支结果,否则只用指定分支做快速推理 |
2. 多分支卷积:trident_conv.py
TridentConv是尺度感知能力的底层载体。构造时需要满足约束:num_branch == len(paddings) == len(dilations)(整数输入会被自动广播为分支长度)。它对外表现为一个nn.Module,前向时按"训练/聚合模式"或"快速推理模式"分别执行分支卷积,并可继续接norm与activation逐分支处理。
3. Trident 骨干网络:trident_backbone.py
TridentBottleneckBlock:标准 ResNet Bottleneck 的 Trident 变体,其中conv2替换为TridentConv,接收num_branch、dilations、concat_output、test_branch_idx。其forward在训练时把单个输入广播复制为num_branch份并行前向;每个 block 的输出是分支列表。最后一个 Trident block 设置concat_output=True,将各分支输出用torch.cat拼回单一张量,保证与后续 RPN/ROIHeads 接口兼容。make_trident_stage:构造一个 ResNet stage,前面若干 block 并行分支,末尾一个 block 用concat_output=True收敛分支。build_trident_resnet_backbone:通过@BACKBONE_REGISTRY.register()注册为MODEL.BACKBONE.NAME = "build_trident_resnet_backbone"。它复用 Detectron2 标准 ResNet 的 stem 与阶段构建逻辑,但在stage_idx == trident_stage_idx(默认 res4)时切换为TridentBottleneckBlock,并显式断言该阶段不支持可变形卷积(Not support deformable conv in Trident blocks yet)。支持DEPTH为 50/101/152(对应num_blocks_per_stage分别为[3,4,6,3]、[3,4,23,3]、[3,8,36,3]),并遵守FREEZE_AT冻结语义。
4. Trident RPN:trident_rpn.py
TridentRPN继承标准RPN并注册进PROPOSAL_GENERATOR_REGISTRY。关键行为是在训练时把图片与 ground-truth 按分支数复制多份(torch.cat([images.tensor] * num_branch)与gt_instances * num_branch),让每个分支都有完整的训练监督;TridentNet-Fast 推理时num_branch退化为 1。
5. Trident ROIHeads:trident_rcnn.py
提供两种已注册的 ROIHeads:
TridentRes5ROIHeads:C4 架构(Res5 做检测头),对应仓库默认配置;TridentStandardROIHeads:标准 FPN 风格 ROIHeads 的 Trident 变体。
两者在训练时同样复制 targets;推理时则调用merge_branch_instances把多个分支的检测结果合并。合并流程清晰可读:先按分支把同一张图的实例拼接(Instances.cat),再做逐类 NMS(batched_nms),最后按test_topk_per_image截取 Top-K 结果。这就是TEST_BRANCH_IDX = -1时"聚合全部分支结果"的实现路径。
训练:端到端启动 TridentNet-Fast
入口脚本与命令
训练入口是 train_net.py,它是官方检测训练脚本的简化版:setup中先get_cfg(),再调用add_tridentnet_config(cfg)注册 TridentNet 专属配置,随后merge_from_file合并 YAML 配置、merge_from_list合并命令行覆盖项并freeze。脚本通过default_argument_parser解析参数,最终用launch启动分布式训练。Trainer继承DefaultTrainer并覆写build_evaluator,返回COCOEvaluator以在训练中同步评估 COCO 指标。
基本训练命令(<config.yaml>为你的配置文件路径):
python /path/to/detectron2/projects/TridentNet/train_net.py --config-file <config.yaml>以 ResNet-50 骨干、8 卡 GPU 端到端训练为例:
python /path/to/detectron2/projects/TridentNet/train_net.py --config-file configs/tridentnet_fast_R_50_C4_1x.yaml --num-gpus 8对应在本仓库内的实际路径为 configs/tridentnet_fast_R_50_C4_1x.yaml。
配置文件逐段解读
基础配置 Base-TridentNet-Fast-C4.yaml 完整定义了 TridentNet-Fast 的架构与训练调度:
MODEL: META_ARCHITECTURE: "GeneralizedRCNN" BACKBONE: NAME: "build_trident_resnet_backbone" ROI_HEADS: NAME: "TridentRes5ROIHeads" POSITIVE_FRACTION: 0.5 BATCH_SIZE_PER_IMAGE: 128 PROPOSAL_APPEND_GT: False PROPOSAL_GENERATOR: NAME: "TridentRPN" RPN: POST_NMS_TOPK_TRAIN: 500 TRIDENT: NUM_BRANCH: 3 BRANCH_DILATIONS: [1, 2, 3] TEST_BRANCH_IDX: 1 TRIDENT_STAGE: "res4" DATASETS: TRAIN: ("coco_2017_train",) TEST: ("coco_2017_val",) SOLVER: IMS_PER_BATCH: 16 BASE_LR: 0.02 STEPS: (60000, 80000) MAX_ITER: 90000 INPUT: MIN_SIZE_TRAIN: (640, 672, 704, 736, 768, 800) VERSION: 2几个关键点:
- 架构拼装:
META_ARCHITECTURE沿用标准的GeneralizedRCNN,仅替换 backbone(build_trident_resnet_backbone)、RPN(TridentRPN)与 ROIHeads(TridentRes5ROIHeads),说明 TridentNet 是一种"即插即用"的尺度感知改造,不动检测器的整体 Meta Architecture; - ROIHeads:
BATCH_SIZE_PER_IMAGE: 128表示每图采样 128 个 RoI(原 Faster R-CNN C4 常用 512,Trident 用 128 即可,也是"Fast"的体现之一);POSITIVE_FRACTION: 0.5控制正负样本比例;PROPOSAL_APPEND_GT: False; - RPN:
POST_NMS_TOPK_TRAIN: 500训练阶段 NMS 后保留 500 个 proposal; - TRIDENT 段:3 分支、空洞率
[1, 2, 3]、Res4 阶段改造、推理用中间分支(索引 1); - 训练调度:
IMS_PER_BATCH: 16(8 卡 × 每卡 2 张)、BASE_LR: 0.02标准线性缩放、STEPS: (60000, 80000)学习率阶梯下降、MAX_ITER: 90000即 1x 训练计划(batch size 16 下的 90k iter ≈ 12 epoch); - 多尺度训练:
MIN_SIZE_TRAIN: (640, 672, 704, 736, 768, 800)在 6 个尺度间随机采样短边,进一步强化对尺度变化的鲁棒性。
模型专用配置通过_BASE_继承基础配置并叠加差异项,例如 R50 1x 配置:
_BASE_: "Base-TridentNet-Fast-C4.yaml" MODEL: WEIGHTS: "detectron2://ImageNetPretrained/MSRA/R-50.pkl" MASK_ON: False RESNETS: DEPTH: 50而 R101 3x 配置(tridentnet_fast_R_101_C4_3x.yaml)仅额外调整预训练权重为 R-101、DEPTH: 101,并把STEPS拉长到(210000, 250000)、MAX_ITER到270000以对应 3x 训练计划。仓库还提供了 tridentnet_fast_R_50_C4_3x.yaml 供 R50 长计划使用。
评估:加载权重执行 COCO 评测
评估与训练共用同一入口,只需增加--eval-only并指定权重文件:
python /path/to/detectron2/projects/TridentNet/train_net.py --config-file configs/tridentnet_fast_R_50_C4_1x.yaml --eval-only MODEL.WEIGHTS model.pth执行流程(对应 train_net.py 的main):eval_only分支中先Trainer.build_model(cfg)构建模型,再用DetectionCheckpointer按cfg.MODEL.WEIGHTS加载(支持resume语义),最后Trainer.test(cfg, model)调用 COCOEvaluator 输出 mAP 指标。推理阶段默认只走TEST_BRANCH_IDX=1(中间分支)这一路分支计算,因此评估成本与普通 Faster R-CNN 一致。
若想以聚合模式评估,将TEST_BRANCH_IDX设为-1即可,此时推理会走TridentRes5ROIHeads中merge_branch_instances的多分支 NMS 合并路径。
MS-COCO 上的实验结果
仓库 README 给出了 Detectron2 实现下 TridentNet-Fast 与 Faster R-CNN 在 COCO 上的对比(COCO 2017 val),完整继承如下:
| Model | Backbone | Head | lr sched | AP | AP50 | AP75 | APs | APm | APl |
|---|---|---|---|---|---|---|---|---|---|
| Faster | R50-C4 | C5-512ROI | 1X | 35.7 | 56.1 | 38.0 | 19.2 | 40.9 | 48.7 |
| TridentFast | R50-C4 | C5-128ROI | 1X | 38.0 | 58.1 | 40.8 | 19.5 | 42.2 | 54.6 |
| Faster | R50-C4 | C5-512ROI | 3X | 38.4 | 58.7 | 41.3 | 20.7 | 42.7 | 53.1 |
| TridentFast | R50-C4 | C5-128ROI | 3X | 40.6 | 60.8 | 43.6 | 23.4 | 44.7 | 57.1 |
| Faster | R101-C4 | C5-512ROI | 3X | 41.1 | 61.4 | 44.0 | 22.2 | 45.5 | 55.9 |
| TridentFast | R101-C4 | C5-128ROI | 3X | 43.6 | 63.4 | 47.0 | 24.3 | 47.8 | 60.0 |
两组关键对比可以清晰看到收益:
- 同骨干、同训练计划下:R50-C4 1X,TridentFast 以 128 个 RoI 的检测头做到 AP 38.0,比 512 个 RoI 的 Faster R-CNN(35.7)高出 2.3 个点,其中大目标(APl)提升最显著(48.7 → 54.6,+5.9);
- 尺度维度放大差距:3X 计划下 R101 骨干,TridentFast AP 达到 43.6,超过同骨干 Faster 的 41.1,AP50 更是达到 63.4。这正印证了"多分支不同感受野、各司其职"对尺度覆盖(尤其大目标)的建模优势。
需要说明的是:以上指标为原项目 README 记录的历史评测结果,在具体硬件与框架版本下复现时可能存在小幅浮动。
引用方式
若你在工作中使用或对比 TridentNet,可沿用原 README 提供的 BibTeX 条目:
@InProceedings{li2019scale, title={Scale-Aware Trident Networks for Object Detection}, author={Li, Yanghao and Chen, Yuntao and Wang, Naiyan and Zhang, Zhaoxiang}, journal={The International Conference on Computer Vision (ICCV)}, year={2019} }在本仓库中的定位
在 IDM-VTON(ECCV 2024 虚拟试穿)仓库中,TridentNet 位于 preprocess/humanparsing/mhp_extension/detectron2 下的 projects 目录,与 DensePose、PointRend、TensorMask 并列,作为 humanparsing 预处理链所依赖的 Detectron2 二次分发版本中的示例检测项目存在。它服务于该子仓库对 Detectron2 代码生态的完整移植:训练脚本通过add_tridentnet_config在get_cfg()基础上扩展配置,这一"以 project 形式扩展核心框架"的组织方式,也正是读者在阅读其他项目(如人类解析网络中的 mask R-CNN 微调配置 parsing_finetune_cihp.yaml 同目录体系)时可以复用的模式。
综上,TridentNet 在 Detectron2 中的实现提供了三个可迁移的工程经验:共享权重 + 多空洞率并行分支让尺度感知成为 backbone 层的标准改造;训练全分支、推理单分支的 Fast 策略实现了"零额外成本换精度";以 config + registry + project 组织的扩展方式让新检测范式可以无缝嵌入成熟的 Meta Architecture 生态。
- 计算机视觉
- 深度学习
- 媒体生成
【免费下载链接】IDM-VTON
[ECCV2024] IDM-VTON : Improving Diffusion Models for Authentic Virtual Try-on in the Wild
相关推荐
WinUI 控件库快速上手:5分钟跑通第一个 Windows 应用界面
WinUI 控件库快速上手:5分钟跑通第一个 Windows 应用界面 WinUI Microsoft.UI.Xaml 是 Windows 的现代 UI 控件库
前端UI组件桌面应用Detectron2 中的 MViTv2 检测实战:多尺度视觉 Transformer 的配置、训练与评估
Detectron2 中的 MViTv2 检测实战:多尺度视觉 Transformer 的配置、训练与评估 MViTv2(Improved Multiscale
人工智能计算机视觉深度学习机器学习基于 Detectron2 的 TridentNet-Fast 目标检测实现:源码剖析、配置详解与训练评估实战
基于 Detectron2 的 TridentNet Fast 目标检测实现:源码剖析、配置详解与训练评估实战 TridentNet(Scale Aware T
人工智能计算机视觉媒体生成AI 应用
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考