☰
MMagic 人脸超分 DIC 算法全解析:迭代协作恢复机制、配置详解与训练实战
2026/9/28 6:32:48 网站建设 项目流程
  • 媒体生成
  • 计算机视觉
  • 深度学习
  • 人工智能
  • 大模型

【免费下载链接】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
点击查看免费下载

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_channels3输入图像通道数(RGB)
out_channels3输出图像通道数
mid_channels64主干网络中间特征通道数(官方配置为 48)
num_blocks6反馈块中上下采样模块数量
hg_mid_channels256Hourglass 中间通道数
hg_num_keypoints68关键点数(68 点人脸 landmark)
num_steps4迭代步数
upscale_factor8上采样倍数
detach_attentionFalse热图是否从当前计算图分离
prelu_init0.2PReLU 的初始斜率
num_heatmaps5融合模块使用的热图组数
num_fusion_blocks7注意力融合模块中残差块数量

前向流程(见forward)清晰地体现了迭代协作:

  1. 输入先经过interpolate双线性插值到 128×128 作为全局残差基线;
  2. conv_first(卷积 + PReLU + PixelShuffle(2))做浅层特征提取;
  3. 第 0 步由FeedbackBlockCustom输出初始 SR 特征;后续每一步由FeedbackBlockHeatmapAttention结合上一步的关键点热图输出特征;
  4. conv_last(转置卷积 + 卷积)得到当前步的 SR 图像,与全局残差相加;
  5. 该 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)按如下顺序构建样本:

  1. LoadImageFromFile:读取 GT 图像(color_type='color',RGB 通道序,cv2 解码);
  2. Resize到 (128, 128)(bicubic,pillow 后端);
  3. Resize缩放 1/8(保持宽高比),输出键为img,即生成 16×16 的低分辨率输入;
  4. GenerateFacialHeatmap:基于 128 尺寸 GT 图生成 32×32 关键点热图(sigma=1.0);
  5. 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_iters150000500000GAN 训练更长
LR milestones[10000, 20000, 40000, 80000][100000, 200000, 300000, 400000]对应更长训练周期
验证间隔20005000—

两套配置均使用MMSeparateDistributedDataParallel作为model_wrapper_cfg(生成器与判别器分离同步),scale = 8表明任务为 8 倍超分。

八、实验结果与模型仓库

在 RGB 通道上评估,评估前裁剪每个边界的scale像素,指标为PSNR / SSIM。需要注意:dic_gan_x8c48b6_g4_150k_CelebAHQ的日志中 DICGAN 仅在 CelebA-HQ 测试集前 9 张图上验证,因此下表中PSNR/SSIM与日志数据不同。

模型数据集scalePSNRSSIM训练资源下载
dic_x8c48b6_g4_150k_CelebAHQCelebAHQx825.23190.74224 (Tesla PG503-216)model | log
dic_gan_x8c48b6_g4_500k_CelebAHQCelebAHQx823.62410.67214 (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.

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

相关推荐

上一篇:FGO智能自动化终极指南:告别重复操作,解放双手的游戏脚本神器
下一篇:OpenCompass项目中的主观评估技术指南

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询