- 媒体生成
- 计算机视觉
- 深度学习
- 人工智能
- 大模型
【免费下载链接】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.
DIC(Deep Face Super-Resolution with Iterative Collaboration)是 2020 年 CVPR 上提出的人脸超分辨率(Face Super-Resolution, FSR)算法,其核心思想是让"图像恢复"与"关键点(Landmark)估计"两个循环网络在多次迭代中互相促进、逐步提升。本文以 MMagic 仓库中 configs/dic/README.md 为主体,结合 mmagic/models/editors/dic 下的完整源码与两套官方配置文件,系统讲解 DIC 的算法原理、网络结构、损失设计、配置参数与训练/测试实战命令,帮助读者在 MMagic 中直接复现人脸 8 倍超分。
一、算法背景与核心思想
论文:Deep Face Super-Resolution with Iterative Collaboration between Attentive Recovery and Landmark Estimation,CVPR 2020任务:图像超分辨率(Image Super-Resolution),模型仓库位于 configs/dic。
近年的深度学习方法借助人脸先验(facial priors)在严重退化人脸图像的超分上取得了不错的效果,但已有方法对先验的利用并不充分:landmark 和 component map 等面部先验通常由低分辨率或粗超分图像估计而来,先验本身不够准确,进而影响了恢复性能。DIC 正是针对这一问题提出的解决方案:
- 两个循环网络迭代协作:一个分支负责面部图像恢复(recovery),另一个分支负责关键点估计(landmark estimation)。在每个循环步中,恢复分支利用关键点先验生成更高质量的图像,而更清晰的图像反过来又促进更准确的关键点估计,二者在迭代中互相增强(iterative information interaction)。
- 新的注意力融合模块(Attentive Fusion Module):将面部组件(facial components)分别生成,再以注意力方式聚合,从而强化 landmark 图对恢复过程的引导。
从 MMagic 的模型注册表看,DIC 的实现分布在 mmagic/models/editors/dic 目录下,核心模块包括:顶层模型 dic.py(DIC)、生成器 dic_net.py(DICNet)、关键点分支 feedback_hour_glass.py(FeedbackHourglass)以及判别器 light_cnn.py(LightCNN)。
二、模型顶层设计:DIC 如何组织迭代训练
MMagic 中的DIC模型定义在 mmagic/models/editors/dic/dic.py,它继承了SRGAN,并在此基础上扩展出两个专属损失:align_loss(热图对齐损失)与feature_loss(基于 LightCNN 的特征感知损失):
@MODELS.register_module() class DIC(SRGAN): def __init__(self, generator, pixel_loss, align_loss, discriminator=None, gan_loss=None, feature_loss=None, train_cfg=None, test_cfg=None, init_cfg=None, data_preprocessor=None): ... self.align_loss = MODELS.build(align_loss) self.feature_loss = MODELS.build(feature_loss) if feature_loss else None self.pixel_init = train_cfg.get('pixel_init', 0) if train_cfg else 0关键设计点:
- 多步输出:
forward_tensor中生成器返回sr_list与heatmap_list两个列表(长度等于迭代步数,默认 4 步)。训练时全部返回用于计算各步损失,推理时只取最后一步sr_list[-1]。 - 判别器延迟启动(pixel_init):
if_run_d()要求self.step_counter >= self.pixel_init才运行判别器,即在训练初期(默认前 10000 步)只做像素级恢复与关键点对齐,让生成器先收敛到稳定的重建质量,再引入对抗训练。 - 判别器更新频率(disc_repeat):
train_step中按disc_repeat(默认 2)次更新判别器,且判别器输入使用sr_list[-1].detach()断开梯度。
从 tests/test_models/test_editors/test_dic/test_dic.py 的单元测试可以验证迭代训练的实际行为:两次train_step后返回的日志变量包含loss_pixel_v0~loss_pixel_v3、loss_align_v0~loss_align_v3、loss_feature、loss_gan、loss_d_real、loss_d_fake共 10 项,正好对应 4 个迭代步的逐级像素损失与热图对齐损失,以及对抗训练的判别损失。
三、生成器 DICNet 源码拆解:反馈块与注意力融合
DICNet定义在 mmagic/models/editors/dic/dic_net.py,其构造函数参数即对应配置文件中的可调项:
| 参数 | 默认值 | 说明 |
|---|---|---|
in_channels | 3 | 输入图像通道数(RGB) |
out_channels | 3 | 输出图像通道数 |
mid_channels | 64 | 主干网络中间特征通道数(官方配置为 48) |
num_blocks | 6 | 反馈块中上下采样模块数量 |
hg_mid_channels | 256 | Hourglass 中间通道数 |
hg_num_keypoints | 68 | 关键点数(68 点人脸 landmark) |
num_steps | 4 | 迭代步数 |
upscale_factor | 8 | 上采样倍数 |
detach_attention | False | 热图是否从当前计算图分离 |
prelu_init | 0.2 | PReLU 的初始斜率 |
num_heatmaps | 5 | 融合模块使用的热图组数 |
num_fusion_blocks | 7 | 注意力融合模块中残差块数量 |
前向流程(见forward)清晰地体现了迭代协作:
- 输入先经过
interpolate双线性插值到 128×128 作为全局残差基线; conv_first(卷积 + PReLU + PixelShuffle(2))做浅层特征提取;- 第 0 步由
FeedbackBlockCustom输出初始 SR 特征;后续每一步由FeedbackBlockHeatmapAttention结合上一步的关键点热图输出特征; conv_last(转置卷积 + 卷积)得到当前步的 SR 图像,与全局残差相加;- 该 SR 图像送入
FeedbackHourglass估计 68 点关键点热图,同时返回隐藏状态last_hidden作为下一步的反馈输入。
其中两类核心模块值得展开:
1. 反馈块(FeedbackBlock)。整体结构呈"模块输出回流到自身输入"的闭环(----- Module ----->上方箭头回流)。FeedbackBlockCustom作为第一个反馈块,会直接把输入特征写入last_hidden作为后续反馈;FeedbackBlockHeatmapAttention则在反馈块中嵌入FeatureHeatmapFusingBlock。
2. 注意力融合模块(FeatureHeatmapFusingBlock)。源码位于同一文件的FeatureHeatmapFusingBlock类,它先将特征经 1×1 卷积扩展到num_heatmaps × in_channels通道,再通过GroupResBlock分组卷积(groups=num_heatmaps)让每组特征独立对应一个面部组件,最后对热图做softmax得到注意力权重,逐通道加权求和完成"组件分别生成、注意力聚合":
attention = nn.functional.softmax(heatmap, dim=1) feature = feature.view(batch_size, self.num_heatmaps, -1, w, h) * attention.unsqueeze(2) feature = feature.sum(1)四、关键点估计分支:FeedbackHourglass 与五组面部热图
FeedbackHourglass定义在 mmagic/models/editors/dic/feedback_hour_glass.py,采用递归式 Hourglass 网络(深度 4,Hourglass(depth-1)递归构造)估计 68 点人脸关键点,并将网络的中间隐藏状态作为反馈传入下一步,形成"预处理 → Hourglass → 反馈"的闭环。
更关键的是reduce_to_five_heatmaps函数:它将 68 点(或 Helen 数据集的 194 点)关键点热图归约为 5 组组件热图,与num_heatmaps=5严格对应:
- 左眼(left eye)
- 右眼(right eye)
- 鼻子(nose)
- 嘴(mouse)
- 人脸轮廓(face silhouette)
归约时先按通道最大值归一化(clamp_min_(0.05)防止除零),再按各组关键点的热图求和。该函数同时支持detach开关:开启后热图从计算图中分离,阻断关键点分支梯度回流到恢复分支,可作为训练策略的调节项。
五、判别器与特征损失:基于 LightCNN 的对抗与感知约束
DICGAN(带 GAN 的 DIC 变体)的判别器LightCNN定义在 mmagic/models/editors/dic/light_cnn.py,输入尺寸 128×128。其基础模块MaxFeature采用"双通道特征取最大"的 max-feature 策略:卷积/线性层输出2 * out_channels通道,再按通道分成两份取逐元素最大值,形成非线性特征选择。
特征损失LightCNNFeatureLoss定义在 mmagic/models/losses/feature_loss.py,它加载预训练的light_cnn_feature.pth(冻结参数、requires_grad_(False))作为特征提取器,在特征空间计算预测图与 GT 的 L1(或 MSE)距离:
pred_feature = self.model(pred) gt_feature = self.model(gt).detach() feature_loss = self.criterion(pred_feature, gt_feature)其中 GT 特征做了detach(),只引导生成器向 GT 特征对齐。
六、数据管线:关键点热图从哪来
训练时所需的 landmark 热图由数据增强变换GenerateFacialHeatmap生成,实现在 mmagic/datasets/transforms/generate_assistant.py:
- 依赖
face-alignment库(源码中assert has_face_alignment明确要求安装),在 CPU 上运行FaceAlignment(LandmarksType._2D)检测 2D 关键点; - 将 128×128 原图上检测到的关键点按
size_ratio缩放到热图尺寸(训练配置为 32×32),并以sigma(默认 1.0)为高斯半径生成逐关键点热图; use_cache=True时以{image_key}_{key}为键缓存热图,避免重复检测。
训练管线(见 dic_x8c48b6_4xb2-150k_celeba-hq.py 的train_pipeline)按如下顺序构建样本:
LoadImageFromFile:读取 GT 图像(color_type='color',RGB 通道序,cv2 解码);Resize到 (128, 128)(bicubic,pillow 后端);Resize缩放 1/8(保持宽高比),输出键为img,即生成 16×16 的低分辨率输入;GenerateFacialHeatmap:基于 128 尺寸 GT 图生成 32×32 关键点热图(sigma=1.0);PackInputs打包。
验证/测试管线与训练管线一致但不生成热图(推理不需要 landmark 监督),另有一套inference_pipeline供推理使用。数据集使用BasicImageDataset,data_root='data',训练集为CelebA-HQ/train_256/all_256,验证/测试集为test_256/all_256。
七、两套官方配置全面解析
DIC 在 MMagic 中提供两套配置,继承关系为:dic_gan-x8c48b6_4xb2-500k_celeba-hq.py以dic_x8c48b6_4xb2-150k_celeba-hq.py为_base_,后者再继承../_base_/default_runtime.py。
7.1 纯重建版:dic_x8c48b6_4xb2-150k_celeba-hq.py
配置文件 只包含生成器与两类损失,无对抗组件:
model = dict( type='DIC', generator=dict( type='DICNet', in_channels=3, out_channels=3, mid_channels=48), pixel_loss=dict(type='L1Loss', loss_weight=1.0, reduction='mean'), align_loss=dict(type='MSELoss', loss_weight=0.1, reduction='mean'), train_cfg=dict(), test_cfg=dict(), data_preprocessor=dict( type='DataPreprocessor', mean=[129.795, 108.12, 96.39], std=[255, 255, 255], ))训练配置要点:
train_cfg:IterBasedTrainLoop,max_iters=150_000,每 2000 步验证一次;- 优化器:
MultiOptimWrapperConstructor+OptimWrapper,生成器Adam(lr=1e-4); - 学习率:
MultiStepLR,milestones=[10000, 20000, 40000, 80000],gamma=0.5; - 评估器:
MAE+PSNR/SSIM(crop_border=scale,即评估前裁剪边界 8 像素); default_hooks:每 2000 步保存 checkpoint(含优化器状态),日志每 100 步输出一次。
7.2 GAN 版:dic_gan-x8c48b6_4xb2-500k_celeba-hq.py
配置文件 在基类配置之上叠加判别器、GAN 损失与特征损失:
model = dict( type='DIC', generator=dict( type='DICNet', in_channels=3, out_channels=3, mid_channels=48), discriminator=dict(type='LightCNN', in_channels=3), pixel_loss=dict(type='L1Loss', loss_weight=1.0, reduction='mean'), align_loss=dict(type='MSELoss', loss_weight=0.1, reduction='mean'), feature_loss=dict( type='LightCNNFeatureLoss', pretrained=pretrained_light_cnn, loss_weight=0.1, criterion='l1'), gan_loss=dict( type='GANLoss', gan_type='vanilla', loss_weight=0.005, real_label_val=1.0, fake_label_val=0), train_cfg=dict(pixel_init=10000, disc_repeat=2), ... )与纯重建版相比的关键差异:
| 配置项 | 纯重建版 | GAN 版 | 说明 |
|---|---|---|---|
pixel_init | 无 | 10000 | 前 10000 步不训练判别器 |
disc_repeat | 无 | 2 | 每轮 G 步后判别器更新 2 次 |
feature_loss | 无 | LightCNNFeatureLoss(L1,权重 0.1) | 需下载预训练light_cnn_feature.pth |
gan_loss | 无 | vanilla GAN,权重 0.005 | 对抗损失 |
| 优化器 | 仅生成器 Adam lr=1e-4 | 生成器 1e-4 / 判别器 1e-5 | 判别器学习率低一个数量级 |
max_iters | 150000 | 500000 | GAN 训练更长 |
| LR milestones | [10000, 20000, 40000, 80000] | [100000, 200000, 300000, 400000] | 对应更长训练周期 |
| 验证间隔 | 2000 | 5000 | — |
两套配置均使用MMSeparateDistributedDataParallel作为model_wrapper_cfg(生成器与判别器分离同步),scale = 8表明任务为 8 倍超分。
八、实验结果与模型仓库
在 RGB 通道上评估,评估前裁剪每个边界的scale像素,指标为PSNR / SSIM。需要注意:dic_gan_x8c48b6_g4_150k_CelebAHQ的日志中 DICGAN 仅在 CelebA-HQ 测试集前 9 张图上验证,因此下表中PSNR/SSIM与日志数据不同。
| 模型 | 数据集 | scale | PSNR | SSIM | 训练资源 | 下载 |
|---|---|---|---|---|---|---|
| dic_x8c48b6_g4_150k_CelebAHQ | CelebAHQ | x8 | 25.2319 | 0.7422 | 4 (Tesla PG503-216) | model | log |
| dic_gan_x8c48b6_g4_500k_CelebAHQ | CelebAHQ | x8 | 23.6241 | 0.6721 | 4 (Tesla PG503-216) | model | log |
对应的模型元信息(配置路径、权重地址、指标)统一登记在 configs/dic/metafile.yml,MMagic 的模型索引系统可据此自动解析与下载权重。
九、快速开始:训练与测试
训练:可使用 CPU、单 GPU 或多 GPU 训练(以 GAN 版配置为例):
# CPU 训练 CUDA_VISIBLE_DEVICES=-1 python tools/train.py configs/dic/dic_gan-x8c48b6_4xb2-500k_celeba-hq.py # 单 GPU 训练 python tools/train.py configs/dic/dic_gan-x8c48b6_4xb2-500k_celeba-hq.py # 多 GPU 训练 ./tools/dist_train.sh configs/dic/dic_gan-x8c48b6_4xb2-500k_celeba-hq.py 8测试:测试命令需要传入预训练权重路径(可先下载上表 model 链接中的权重):
# CPU 测试 CUDA_VISIBLE_DEVICES=-1 python tools/test.py configs/dic/dic_gan-x8c48b6_4xb2-500k_celeba-hq.py https://download.openmmlab.com/mmediting/restorers/dic/dic_gan_x8c48b6_g4_500k_CelebAHQ_20210625-3b89a358.pth # 单 GPU 测试 python tools/test.py configs/dic/dic_gan-x8c48b6_4xb2-500k_celeba-hq.py https://download.openmmlab.com/mmediting/restorers/dic/dic_gan_x8c48b6_g4_500k_CelebAHQ_20210625-3b89a358.pth # 多 GPU 测试 ./tools/dist_test.sh configs/dic/dic_gan-x8c48b6_4xb2-500k_celeba-hq.py https://download.openmmlab.com/mmediting/restorers/dic/dic_gan_x8c48b6_g4_500k_CelebAHQ_20210625-3b89a358.pth 8更多训练与测试细节可参考 docs/en/user_guides/train_test.md(中文版见 docs/zh_cn/user_guides/train_test.md)中Train a model与Test a pre-trained model两部分。运行训练/测试时,请确保已按仓库说明完成 MMagic 安装,并提前准备好 CelebA-HQ 数据集(按data/CelebA-HQ/train_256/all_256与test_256/all_256目录组织);训练 GAN 版还需满足face-alignment依赖(用于热图生成)。
十、引用
若在研究中使用了 DIC,请按如下方式引用:
@inproceedings{ma2020deep, title={Deep face super-resolution with iterative collaboration between attentive recovery and landmark estimation}, author={Ma, Cheng and Jiang, Zhenyu and Rao, Yongming and Lu, Jiwen and Zhou, Jie}, booktitle={Proceedings of the IEEE/CVF conference on computer vision and pattern recognition}, pages={5569--5578}, year={2020} }小结
DIC 以"恢复分支与关键点分支迭代协作 + 组件级注意力融合"为核心,在 8 倍人脸超分任务上兼顾了重建精度与先验引导的稳定性。在 MMagic 中,读者既可以通过 dic_x8c48b6_4xb2-150k_celeba-hq.py 快速复现纯重建版本,也可以通过 dic_gan-x8c48b6_4xb2-500k_celeba-hq.py 体验完整的对抗训练流程;结合 mmagic/models/editors/dic 下的源码,可以深入理解迭代协作、反馈机制与注意力融合的每一步实现细节。
- 媒体生成
- 计算机视觉
- 深度学习
- 人工智能
- 大模型
【免费下载链接】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.
相关推荐
MMagic 中的 DIC:基于迭代协作机制的人脸 8 倍超分辨率算法实战指南
MMagic 中的 DIC:基于迭代协作机制的人脸 8 倍超分辨率算法实战指南 本文聚焦 OpenMMLab 生成式视觉工具箱 MMagic 中的人脸超分辨率算
媒体生成计算机视觉深度学习人工智能大模型MMagic 中的 DeepFillv1 图像修复算法:配置解析、两阶段训练原理与完整实战指南
MMagic 中的 DeepFillv1 图像修复算法:配置解析、两阶段训练原理与完整实战指南 本文导读 :DeepFillv1 是 2018 年 CVPR 上
媒体生成计算机视觉深度学习人工智能大模型MMagic 图像抠图实战:GCA(Guided Contextual Attention)算法原理、配置文件与训练测试全解析
MMagic 图像抠图实战:GCA(Guided Contextual Attention)算法原理、配置文件与训练测试全解析 GCA(Guided Conte
媒体生成计算机视觉深度学习人工智能大模型
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考