• 媒体生成
  • 计算机视觉
  • 深度学习
  • 人工智能
  • 大模型

【免费下载链接】mmagic

OpenMMLab Multimodal Advanced, Generative, and Intelligent Creation Toolbox. Unlock the magic 🪄: Generative-AI (AIGC), easy-to-use APIs, awsome model zoo, diffusion models, for text-to-image generation, image/video restoration/enhancement, etc.

项目地址: https://gitcode.com/gh_mirrors/mm/mmagic
点击查看 免费下载

导读

本文围绕 OpenMMLab MMagic 工具箱中 TOFlow 算法(Video Enhancement with Task-Oriented Flow,IJCV 2019)展开,系统讲解其"任务导向光流"的核心思想、在 MMagic 中针对视频插帧(Video Frame Interpolation)与视频超分辨率(Video Super-Resolution)两任务的具体网络实现、预训练 SPyNet 权重加载机制、Vimeo-90K/Vid4 评测配置,以及完整的训练与测试命令行操作。读完本文,你将掌握 TOFlow 在 MMagic 中的模型结构(TOFlowVFINet / TOFlowVSRNet)、配置文件中关键参数的含义,以及如何基于 tools/train.py 与 tools/test.py 一键复现或复测官方模型。

1. TOFlow 算法背景与核心思想

TOFlow(Task-Oriented Flow)发表于 IJCV 2019,论文题为 Video Enhancement with Task-Oriented Flow。传统视频增强算法通常依赖光流(optical flow)来对齐视频帧序列中的相邻帧,但精确的光流估计本身是困难的,而且对特定视频处理任务而言,通用光流往往不是最优的运动表示。

TOFlow 的核心思路是:设计一个包含可训练运动估计组件与视频处理组件的神经网络,两者联合训练,从而在自监督、任务特定的方式下学习出"任务导向光流"(task-oriented flow)——即为插帧、去噪/去块、超分等具体任务量身定制的运动表示。为支撑低层视频处理研究,论文作者还构建了 Vimeo-90K 大规模高质量视频数据集。在帧插值、视频去噪/去块、视频超分三个任务上,TOFlow 在标准基准以及 Vimeo-90K 上均优于传统光流方法。

在 MMagic 中,TOFlow 被同时用于两类任务(见 configs/tof/README.md):

  • Video Interpolation(视频插帧):使用 TOFlowVFINet,基于相邻两帧插值出中间帧;
  • Video Super-Resolution(视频超分):使用 TOFlowVSRNet,对输入的低分辨率视频帧序列进行重建。

2. 模型结构与源码实现

TOFlow 的两个网络实现位于 mmagic/models/editors/tof/,分别对应插帧与超分两个任务,但共享同一套 SPyNet 光流金字塔子网络。

2.1 视频插帧网络 TOFlowVFINet

TOFlowVFINet 定义于 mmagic/models/editors/tof/tof_vfi_net.py,结构上由两部分组成:

  • 运动估计模块 self.spynet:SPyNet 实例,其 flow_cfg 支持传入 norm_cfg(归一化配置,插帧任务中为 None,即无 BN)与 pretrained(预训练权重路径);
  • 重建模块 self.resnet:三层 ToFResBlock(6→64→64→3 通道,卷积核 9/1/1)。

前向过程(输入 (b, 2, 3, h, w),输出插值帧 (b, 3, h, w))为:

flow_10 = self.spynet(imgs[:, 0], imgs[:, 1]).permute(0, 2, 3, 1)  # 前向流
flow_01 = self.spynet(imgs[:, 1], imgs[:, 0]).permute(0, 2, 3, 1)  # 后向流
warp_frame0 = flow_warp(imgs[:, 0], flow_01 / 2)                   # 用半程流对齐
warp_frame1 = flow_warp(imgs[:, 1], flow_10 / 2)
output = self.resnet(torch.stack([warp_frame0, warp_frame1], dim=1))

即先估计两个方向的光流,取半程流(flow / 2)将前后帧分别 warp 到中间时刻,再由残差重建模块融合输出。ToFResBlock 内部还保留了残差连接:输出 = 重建分支结果 + 两帧均值。

2.2 视频超分网络 TOFlowVSRNet

TOFlowVSRNet 定义于 mmagic/models/editors/tof/tof_vsr_net.py。与插帧版本的关键差异包括:

  • 输入为 7 帧 (b, 7, 3, h, w),其中低分辨率帧已预上采样到与 GT 相同的尺寸(论文设定);
  • 通过 adapt_official_weights 参数适配官方权重的帧序:官方实现以第 0 帧为参考帧,置为 True 时会将输入重排为 [3, 0, 1, 2, 4, 5, 6],ref_idx 取 0;置为 False(从头训练)时 ref_idx 取 3,即中间帧为参考帧;
  • 重建模块为 4 层卷积(3*7 → 64 → 64 → 64 → 3,卷积核 9/9/1/1),并在输出处加回参考帧(残差学习)。

前向时,对除参考帧外的其余 6 帧逐一用 SPyNet 估计到参考帧的光流并 warp 对齐,最后堆叠 7 个对齐帧送入重建卷积。

2.3 共享的 SPyNet 子网络

两个网络内的 SPyNet(见 tof_vfi_net.py 与 tof_vsr_net.py)是专为 TOFlow 改造的版本,与通用 SPyNet 有两点区别:

  1. TOFlow 论文中的基本模块包含 BatchNorm(超分版本每个 BasicModule 内部 4 个 ConvModule 均带 norm_cfg=dict(type='BN'));
  2. 归一化/反归一化不在 SPyNet 内部完成,而是交由 TOFlow 网络整体处理(网络注册了 mean / std buffer,默认 RGB 均值 [0.485, 0.456, 0.406]、方差 [0.229, 0.224, 0.225])。

其前向采用经典的空间金字塔 + 由粗到细策略:对参考帧与支撑帧连续做 3 次 2× 平均池化下采样(共 4 个尺度),在 h//16 × w//16 的最低分辨率上以零初始化光流,然后逐层用双线性插值将粗尺度光流放大 2 倍,与当前尺度帧拼接后送入 BasicModule 估计残差流并累加,最终得到全分辨率光流。每个 BasicModule 的输入为 8 通道:参考帧(3)+ 邻居帧(3)+ 初始光流(2),输出 2 通道光流残差。

光流 warp 操作由通用工具 mmagic/models/utils/flow_warp.py 提供,其基于 grid_sample 实现双线性可微 warp,且要求输入张量与光流张量空间尺寸一致。

2.4 顶层模型封装

单元测试 tests/test_models/test_editors/test_tof/test_tof_vfi_net.py 与 tests/test_models/test_editors/test_tof/test_tof_vsr_net.py 分别验证了:输入 (1, 2, 3, 256, 256) 时 TOFlowVFINet 输出 (1, 3, 256, 256);TOFlowVSRNet 在 adapt_official_weights 为 True/False 两种模式下,输入 (2, 7, 3, 16, 16) 均输出 (2, 3, 16, 16)。

3. 配置文件深度解读

3.1 插帧训练配置(base_tof.py + 具体 SPyNet 变体)

所有 5 个插帧配置共享基类 configs/base/models/base_tof.py,其关键内容:

  • 数据集:BasicFramesDataset,数据根目录 data/vimeo_triplet;训练使用 tri_trainlist.txt,验证/测试使用 tri_testlist.txt,load_frames_list 指定输入帧 im1.png、im3.png,GT 为中间帧 im2.png;
  • 训练循环:IterBasedTrainLoop,max_iters=1_000_000,val_interval=5000(5000 iters ≈ 1 epoch);
  • 优化器:Adam,lr=5e-5,betas=(0.9, 0.99),weight_decay=1e-4;
  • 学习率策略:MultiStepLR(by_epoch=False),gamma=0.5,里程碑为 [200000, 400000, 600000, 800000];
  • 评测指标:MAE / PSNR / SSIM;
  • Hook:每 5000 iters 存一次 checkpoint(含优化器状态),日志间隔 100 iters。

具体变体配置(如 tof_spynet-chair-wobn_1xb1_vimeo90k-triplet.py)只做两件事:指定实验名 experiment_name / work_dir / save_dir,以及通过 load_pretrained_spynet 加载对应数据集的预训练 SPyNet 权重。模型部分统一为:

model = dict(
    type='BasicInterpolator',
    generator=dict(
        type='TOFlowVFINet',
        flow_cfg=dict(norm_cfg=None, pretrained=load_pretrained_spynet)),
    pixel_loss=dict(type='CharbonnierLoss', loss_weight=1.0, reduction='mean'),
    train_cfg=dict(),
    test_cfg=dict(),
    required_frames=2,
    step_frames=1,
    init_cfg=None,
    data_preprocessor=dict(
        type='DataPreprocessor',
        mean=[0.485 * 255, 0.456 * 255, 0.406 * 255],
        std=[0.229 * 255, 0.224 * 255, 0.225 * 255],
        pad_size_divisor=16,
        pad_mode='reflect',
    ))

要点说明:

  • flow_cfg.norm_cfg=None 表示插帧版 SPyNet 不含 BN(因为 batch_size=1,与官方 pytoflow 实现保持一致);其余 4 个配置分别把 pretrained 换成 kitti / sintel-clean / sintel-final / pytoflow 对应的权重;
  • data_preprocessor 中均值/方差乘了 255,是因为输入图像为 0–255 范围,需按 ImageNet 归一化统计做标准化;
  • pad_size_divisor=16 与 pad_mode='reflect' 确保输入尺寸被 padding 到 16 的倍数(匹配 SPyNet 4 层 2× 下采样),便于任意尺寸推理。

3.2 超分官方模型配置(tof_x4_official_vimeo90k.py)

tof_x4_official_vimeo90k.py 为 4× 视频超分官方模型,仅用于测试官方权重(配置文件注释明确说明 only testing is supported):

model = dict(
    type='EDVR',  # use the shared model with EDVR
    generator=dict(type='TOFlowVSRNet', adapt_official_weights=True),
    pixel_loss=dict(type='CharbonnierLoss', loss_weight=1.0, reduction='sum'),
    data_preprocessor=dict(
        type='DataPreprocessor',
        mean=[0.485 * 255, 0.456 * 255, 0.406 * 255],
        std=[0.229 * 255, 0.224 * 255, 0.225 * 255],
    ))

其验证数据为 Vid4:data_root='data/Vid4',低分辨率输入前缀 BIx4up_direct(已 4× 双线性上采样)、GT 前缀 GT,标注文件 meta_info_Vid4_GT.txt,num_input_frames=7;验证 pipeline 使用 GenerateFrameIndiceswithPadding(padding='reflection_circle',首尾帧反射补帧以凑足 7 帧窗口)。注意该配置中 test_dataloader 与 test_cfg 被注释为 TODO(测试数据尚未上传),官方评测目前以 val loop 形式在 Vid4 上执行。

4. 模型库与评测结果

4.1 Vimeo-90K Triplet 上的插帧结果(PSNR)

评测于 Vimeo90k-triplet(RGB 通道),指标 PSNR,训练资源为 1 块 Tesla PG503-216:

模型Pretrained SPyNetPSNR
tof_vfi_spynet_chair_nobn_1xb1_vimeo90kspynet_chairs_final33.3294
tof_vfi_spynet_kitti_nobn_1xb1_vimeo90kspynet_chairs_final33.3339
tof_vfi_spynet_sintel_clean_nobn_1xb1_vimeo90kspynet_chairs_final33.3170
tof_vfi_spynet_sintel_final_nobn_1xb1_vimeo90kspynet_chairs_final33.3237
tof_vfi_spynet_pytoflow_nobn_1xb1_vimeo90kspynet_chairs_final33.3426

4.2 Vimeo-90K Triplet 上的插帧结果(SSIM)

同一批模型以 SSIM 作为第二指标评测:

模型SSIM
tof_vfi_spynet_chair_nobn_1xb1_vimeo90k0.9465
tof_vfi_spynet_kitti_nobn_1xb1_vimeo90k0.9466
tof_vfi_spynet_sintel_clean_nobn_1xb1_vimeo90k0.9464
tof_vfi_spynet_sintel_final_nobn_1xb1_vimeo90k0.9465
tof_vfi_spynet_pytoflow_nobn_1xb1_vimeo90k0.9467

注:上述预训练 SPyNet 均不含 BN 层(batch_size=1),与官方 pytoflow 实现保持一致。表格中的下载权重与训练日志可通过 README 中的 Download 链接获取,训练时通过 load_pretrained_spynet 指向对应 .pth。

4.3 Vid4 上的 4× 超分结果

评测于 RGB 通道,指标 PSNR / SSIM,数据集 Vid4:

模型数据集任务Vid4 (PSNR / SSIM)
tof_x4_vimeo90k_officialvimeo90kVideo Super-Resolution24.4377 / 0.7433

5. 快速上手:训练与测试命令

本节命令均来自 configs/tof/README.md 的 Quick Start 部分,并补充了任务/资源层面的说明。执行前需先按 MMagic 文档完成环境安装与数据准备(插帧数据放入 data/vimeo_triplet,超分数据放入 data/Vid4)。更完整的训练/测试流程可参阅 docs/en/user_guides/train_test.md。

5.1 训练

TOF 当前仅支持视频插帧任务的训练(超分官方模型仅用于测试)。以下命令以 chair SPyNet 变体为例:

# CPU 训练
CUDA_VISIBLE_DEVICES=-1 python tools/train.py configs/tof/tof_spynet-chair-wobn_1xb1_vimeo90k-triplet.py

# 单卡训练
python tools/train.py configs/tof/tof_spynet-chair-wobn_1xb1_vimeo90k-triplet.py

# 多卡训练(以 8 卡为例)
./tools/dist_train.sh configs/tof/tof_spynet-chair-wobn_1xb1_vimeo90k-triplet.py 8

训练过程中会按 base_tof.py 的调度在 200k/400k/600k/800k iter 各将学习率减半,每 5000 iters 保存一次 checkpoint,总训练 100 万 iter。

5.2 测试

TOF 支持两类任务的测试:视频插帧与视频超分。

任务 1:视频插帧(以 chair SPyNet 变体为例):

# CPU 测试
CUDA_VISIBLE_DEVICES=-1 python tools/test.py configs/tof/tof_spynet-chair-wobn_1xb1_vimeo90k-triplet.py https://download.openmmlab.com/mmediting/video_interpolators/toflow/pretrained_spynet_chair_20220321-4d82e91b.pth

# 单卡测试
python tools/test.py configs/tof/tof_spynet-chair-wobn_1xb1_vimeo90k-triplet.py https://download.openmmlab.com/mmediting/video_interpolators/toflow/pretrained_spynet_chair_20220321-4d82e91b.pth

# 多卡测试(以 8 卡为例)
./tools/dist_test.sh configs/tof/tof_spynet-chair-wobn_1xb1_vimeo90k-triplet.py https://download.openmmlab.com/mmediting/video_interpolators/toflow/pretrained_spynet_chair_20220321-4d82e91b.pth 8

任务 2:视频超分(Vid4 上的 4× 官方模型):

# CPU 测试
CUDA_VISIBLE_DEVICES=-1 python tools/test.py configs/tof/tof_x4_official_vimeo90k.py https://download.openmmlab.com/mmediting/restorers/tof/tof_x4_vimeo90k_official-a569ff50.pth

# 单卡测试
python tools/test.py configs/tof/tof_x4_official_vimeo90k.py https://download.openmmlab.com/mmediting/restorers/tof/tof_x4_vimeo90k_official-a569ff50.pth

# 多卡测试(以 8 卡为例)
./tools/dist_test.sh configs/tof/tof_x4_official_vimeo90k.py https://download.openmmlab.com/mmediting/restorers/tof/tof_x4_vimeo90k_official-a569ff50.pth 8

命令中的第二个参数为权重 URL:插帧任务传入预训练 SPyNet 权重(pretrained_spynet_*),超分任务传入官方完整模型权重(tof_x4_vimeo90k_official-*.pth)。多卡脚本 tools/dist_train.sh 与 tools/dist_test.sh 内部通过 PORT 环境变量与 torch.distributed 启动分布式进程。

6. 引用信息

若在学术工作中使用 TOFlow 或本实现,可按 README 提供的 BibTeX 引用:

@article{xue2019video,
  title={Video enhancement with task-oriented flow},
  author={Xue, Tianfan and Chen, Baian and Wu, Jiajun and Wei, Donglai and Freeman, William T},
  journal={International Journal of Computer Vision},
  volume={127},
  number={8},
  pages={1106--1125},
  year={2019},
  publisher={Springer}
}

小结

TOFlow 在 MMagic 中的落地体现了"任务导向光流"设计思想的可复现性:插帧侧由 TOFlowVFINet 用双向半程光流 + 残差重建完成帧插值;超分侧由 TOFlowVSRNet 用 7 帧窗口 + SPyNet 对齐 + 卷积重建完成 4× 重建。配合 configs/base/models/base_tof.py 的完整训练流水线与 tools/train.py、tools/test.py 的一键命令,研究者既可以快速复现 Vimeo-90K Triplet 插帧实验,也可以直接加载官方权重在 Vid4 上复测超分指标。

  • 媒体生成
  • 计算机视觉
  • 深度学习
  • 人工智能
  • 大模型

【免费下载链接】mmagic

OpenMMLab Multimodal Advanced, Generative, and Intelligent Creation Toolbox. Unlock the magic 🪄: Generative-AI (AIGC), easy-to-use APIs, awsome model zoo, diffusion models, for text-to-image generation, image/video restoration/enhancement, etc.

项目地址: https://gitcode.com/gh_mirrors/mm/mmagic
点击查看 免费下载
Logo

火山引擎视频云技术社区,是面向 AI 音视频开发者的技术交流平台。这里汇聚源自抖音、豆包等亿级 DAU 产品的 RTC、直播、点播、AI 媒体处理、音视频互动技术,提供接入指南、最佳实践、性能调优、场景案例、Demo 代码、开源项目、白皮书和 API 文档。社区汇聚官方工程师与一线开发者,为 AI 视频通话、数字人、AI 视频处理等应用的开发与落地提供技术支持。

更多推荐