☰
mmagic 中 SAGAN 条件生成对抗网络的模型配置、训练与评估实战指南
2026/9/29 6:14:28 网站建设 项目流程
  • 媒体生成
  • 计算机视觉
  • 深度学习
  • 人工智能
  • 大模型

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

本文基于 configs/sagan/README.md 展开,结合 mmagic 仓库中 SAGAN 的模型实现、基础配置与评估工具,系统讲解如何在 mmagic 中复现与使用 SAGAN(Self-Attention Generative Adversarial Networks,ICML'2019)。读完本文,你将掌握 SAGAN 的算法动机、其在 mmagic 中的源码结构与关键超参数、CIFAR10 与 ImageNet 两套官方训练配置的逐项含义、迭代计数换算规则,以及使用 Inception Score(IS)与 Fréchet Inception Distance(FID)评估生成质量的完整流程,并了解如何将 PyTorch-StudioGAN 的预训练权重转换到 mmagic 中使用。

一、SAGAN 算法核心思想

SAGAN 由 Zhang Han 等人发表于 ICML 2019(论文标题Self-attention generative adversarial networks)。其核心动机在于:传统卷积 GAN 在生成高分辨率细节时,仅依赖于低分辨率特征图中空间局部的点,缺乏对远距离依赖关系的建模能力;而 SAGAN 引入注意力机制,使生成器能够利用全部特征位置的线索来生成细节,判别器也能够校验图像中相距较远的区域之间细节是否彼此一致。

除此之外,SAGAN 还借鉴了"生成器条件化(generator conditioning)影响 GAN 性能"的发现,将**谱归一化(Spectral Normalization)**应用于生成器,从而改善训练动态。在原论文中,SAGAN 在极具挑战性的 ImageNet 数据集上将最佳 Inception Score(IS)从 36.8 提升到 52.52,并将 Fréchet Inception Distance(FID)从 27.62 降低到 18.65;注意力层可视化显示,生成器关注的是与物体形状对应的邻域,而非固定形状的局部区域。

在 mmagic 中,SAGAN 被归类于Conditional GANs(条件生成对抗网络)任务,模型集合注册信息与结果元数据可在 configs/sagan/metafile.yml 中查看。

二、mmagic 中的源码实现结构

mmagic 对 SAGAN 的实现位于mmagic/models/editors/sagan/目录,包含四个核心文件:

  • sagan.py:顶层模型类SAGAN,继承自BaseConditionalGAN,负责组织生成器与判别器的训练流程;
  • sagan_generator.py:生成器SNGANGenerator(注册名SAGANGenerator);
  • sagan_discriminator.py:投影判别器ProjDiscriminator;
  • sagan_modules.py:生成器 ResBlock(SNGANGenResBlock)、条件归一化(SNConditionNorm)等基础模块。

从类注释与代码实现可以推断,mmagic 的 SAGAN 实现同时融合了三个相关工作:

组件来源工作作用
SAGANSelf-Attention GAN自注意力长期依赖建模
SNGANGeneratorSpectral Normalization GAN(SNGAN)生成器谱归一化
ProjDiscriminatorcGANs with Projection Discriminator(Proj-GAN)投影式条件判别器

自注意力模块SelfAttentionBlock则复用了 BigGAN 实现中的模块(见 biggan_modules.py 相关定义),在基础配置中通过attention_cfg=dict(type=SelfAttentionBlock)注入。

2.1 顶层模型与损失函数

SAGAN类在 sagan.py 中通过@MODELS.register_module('SNGAN')与@MODELS.register_module()双重注册。其构造函数接受generator、discriminator、data_preprocessor、generator_steps、discriminator_steps、noise_size(默认 128)、num_classes、ema_config等参数。

该实现使用Hinge Loss训练生成器与判别器:

  • 判别器损失(disc_loss):loss_disc_fake = relu(1 + D(fake))的均值加上loss_disc_real = relu(1 - D(real))的均值;
  • 生成器损失(gen_loss):loss_gen = -D(fake).mean()。

在train_discriminator中,真实图像取自data_samples.gt_img,标签取自data_sample_to_label,生成器输出在torch.no_grad()下计算;在train_generator中则直接以噪声noise_fn与随机标签label_fn生成假图并计算对抗损失。

2.2 生成器 SNGANGenerator

SNGANGenerator的关键设计是channels_cfg与blocks_cfg两个可配置项:

  • channels_cfg:在 SNGAN / Proj-GAN 的默认配置中,ResBlock 数量与各层通道数与输出分辨率一一对应。代码内置了_default_channels_cfg字典:
_default_channels_cfg = { 32: [1, 1, 1], 64: [16, 8, 4, 2], 128: [16, 16, 8, 4, 2] }

即只需要给出output_scale,即可自动推导通道结构;也支持用户自定义列表或字典。

  • blocks_cfg:默认使用dict(type='SNGANGenResBlock'),用户可通过MODELS.build机制替换中间块,提高模型泛化性。

生成器的前向流程为:噪声noise(形状(n, noise_size))经noise2feat线性层映射并 reshape 为input_scale × input_scale的特征图,随后依次经过若干SNGANGenResBlock(必要时插入SelfAttentionBlock),最后经to_rgb卷积与Tanh激活输出图像。attention_after_nth_block参数(int 或 int 列表)决定自注意力块插入到第几个 ConvBlock 之后;num_classes=0时条件归一化层会自动退化为无条件版本。

此外,init_weights支持多种初始化风格:STUDIO(Pytorch-StudioGAN,正交初始化)、BIGGAN(xavier_uniform)、SAGAN(官方 TensorFlow 实现)、SNGAN/SNGAN-PROJ/GAN-PROJ(官方 Chainer 实现),对应init_cfg中的type字段。

2.3 谱归一化相关超参数

SNGANGenerator与SNGANGenResBlock提供了若干与谱归一化、归一化稳定性相关的细粒度参数,理解它们对调参很有帮助:

参数默认值含义
with_spectral_normFalse卷积块是否使用谱归一化
with_embedding_spectral_normNone归一化块中 embedding 层是否谱归一化;未指定时跟随with_spectral_norm
sn_style'torch'谱归一化实现风格:torch(PyTorch 官方实现)或ajbrock(BigGAN-PyTorch 实现)
sn_eps1e-12谱归一化操作的 epsilon
norm_eps1e-4条件/非条件归一化层的 epsilon
auto_sync_bnTrue分布式训练时是否将 BatchNorm 转为 SyncBN

三、官方配置逐项解析

configs/sagan/目录共提供 6 个配置文件,其中基础模型配置集中在 mmagic/configs/base/models/sagan/base_sagan_32x32.py 与 mmagic/configs/base/models/sagan/base_sagan_128x128.py。

3.1 基础模型配置(32×32 与 128×128)

以 32×32(CIFAR10)基础配置为例:

model = dict( type=SAGAN, data_preprocessor=dict(type=DataPreprocessor), num_classes=10, generator=dict( type=SNGANGenerator, num_classes=10, output_scale=32, base_channels=256, attention_cfg=dict(type=SelfAttentionBlock), attention_after_nth_block=2, with_spectral_norm=True), discriminator=dict( type=ProjDiscriminator, num_classes=10, input_scale=32, base_channels=128, attention_cfg=dict(type=SelfAttentionBlock), attention_after_nth_block=1, with_spectral_norm=True), generator_steps=1, discriminator_steps=5)

128×128(ImageNet)基础配置与之对应:num_classes=1000、生成器output_scale=128、base_channels=64、attention_after_nth_block=4,判别器input_scale=128、base_channels=64、attention_after_nth_block=1,且generator_steps=1、discriminator_steps=1。可见在 128×128 场景下生成器在第 4 个 ResBlock 后插入自注意力,通道数也从 64 起步(通道倍率遵循_default_channels_cfg[128])。

3.2 CIFAR10 32×32 训练配置

配置文件 sagan_woReLUinplace_lr2e-4-ndisc5-1xb64_cifar10-32x32.py 继承gen_default_runtime.py、cifar10_nopad.py数据集与base_sagan_32x32.py基础模型,其关键训练设置:

  • disc_step = 5:每更新一次生成器前先更新 5 次判别器;
  • init_cfg = dict(type='studio'):采用 Pytorch-StudioGAN 风格的正交初始化;
  • data_preprocessor=dict(output_channel_order='BGR'):CIFAR 图像为 RGB,需转换为 BGR 通道顺序;
  • train_cfg = dict(max_iters=100000 * disc_step):总迭代数 = 100000 × 5;
  • train_dataloader = dict(batch_size=64):单卡 batch size 64;
  • 优化器:生成器与判别器均使用 Adam,lr=0.0002,betas=(0.5, 0.999);
  • VisualizationHook:每 5000 次迭代可视化一次固定输入的生成结果(fixed_input=True,vis_kwargs_list=dict(type='GAN', name='fake_img'))。

3.3 ImageNet 128×128 训练配置

配置文件 sagan_woReLUinplace_Glr1e-4_Dlr4e-4_ndisc1-4xb64_imagenet1k-128x128.py 的关键设置:

  • 生成器与判别器学习率解耦:生成器Adam lr=0.0001、判别器Adam lr=0.0004,两者betas=(0.0, 0.999);
  • discriminator_steps=1(ndisc1):判别器与生成器交替更新;
  • train_cfg = dict(max_iters=1000000, val_interval=10000, dynamic_intervals=[(800000, 4000)]):总迭代 100 万次,每 1 万次迭代验证一次,80 万次迭代后验证间隔动态调整为 4000;
  • train_dataloader = dict(batch_size=64)(注释标注为 4 卡训练,即总 batch size 64×4)。

3.4 BigGAN Schedule 变体配置

配置文件 sagan_woReLUinplace-Glr1e-4_Dlr4e-4_noaug-ndisc1-8xb32-bigGAN-sch_imagenet1k-128x128.py 遵循 BigGAN 官方仓库launch_SAGAN_bz128x2_ema.sh的设置,其注释明确列出 6 点差异:

  1. 谱归一化使用eps=1e-8;
  2. 不使用 SyncBN(auto_sync_bn=False);
  3. 条件归一化(cBN)中的 embedding 层不使用谱归一化(with_embedding_spectral_norm=False);
  4. 在特定迭代开始启用 EMA(ema_config=dict(interval=1, momentum=0.999, start_iter=2000));
  5. 权重初始化使用xavier_uniform(init_cfg = dict(type='BigGAN'));
  6. 不进行数据增强(继承imagenet_noaug_128.py数据集)。

此外,该配置的生成器还设置了norm_eps=1e-5、sn_eps=1e-8,判别器sn_eps=1e-8;可视化 Hook 同时输出 EMA 与原始权重模型的生成结果(sample_model='ema/orig',target_keys=['ema.fake_img', 'orig.fake_img']),评估指标则使用 EMA 模型(sample_model='ema')。

四、评估指标配置

所有训练配置末尾都挂载了 IS 与 FID 两个指标,例如:

inception_pkl = './work_dirs/inception_pkl/cifar10-full.pkl' metrics = [ dict( type='InceptionScore', prefix='IS-50k', fake_nums=50000, inception_style='StyleGAN', sample_model='orig'), dict( type='FrechetInceptionDistance', prefix='FID-Full-50k', fake_nums=50000, inception_style='StyleGAN', inception_pkl=inception_pkl, sample_model='orig') ] default_hooks = dict( checkpoint=dict( save_best=['FID-Full-50k/fid', 'IS-50k/is'], rule=['less', 'greater']))

要点说明:

  • fake_nums=50000:评估时生成 5 万张假图;
  • inception_style='StyleGAN':使用 Tero 的 Inception V3 script module 提取特征(详见 docs/en/user_guides/metrics.md);
  • inception_pkl:FID 需要真实数据集的 Inception 特征统计量,预先保存为 pkl 可避免每次评估重复提取;
  • default_hooks.checkpoint:同时保存 FID 最小与 IS 最大的两个最优 checkpoint(rule=['less', 'greater'])。

关于 Inception V3 与图像缩放方式的选择,mmagic 的指标文档明确指出这两者会显著影响最终 IS 分数,因此强烈推荐使用 Tero 的 script model(加载需要torch >= 1.6),并采用Pillow 后端的 Bicubic 插值进行缩放。对应配置中可通过resize_method与use_pillow_resize设置缩放方式,通过inception_style选择StyleGAN(Tero 模型)或PyTorch(torchvision 实现),在无网络环境下可下载 Inception 权重并通过inception_path指定。

五、迭代计数规则与实验结果

5.1 迭代计数换算

原文档特别强调:mmagic 实现的迭代计数规则与其他代码库不同。若需与其他代码库对齐,可使用如下换算公式:

total_iters (biggan/pytorch studio gan) = our_total_iters / dist_step

其中dist_step即配置中的disc_step(判别器每轮更新次数)。例如 CIFAR10 配置中disc_step=5、max_iters=500000,对应其他代码库的 100000 次迭代。

5.2 官方训练结果

以下是 mmagic 官方在 CIFAR10 与 ImageNet 上训练的 SAGAN 模型结果(模型权重可通过 configs/sagan/metafile.yml 中对应条目的Weights字段获取):

模型数据集Inplace ReLUdist_step总 batch size总迭代数*最佳迭代ISFID
SAGAN-32x32-woInplaceReLU Best ISCIFAR10w/o564×15000004000009.321710.5030
SAGAN-32x32-woInplaceReLU Best FIDCIFAR10w/o564×15000004800009.31749.4252
SAGAN-32x32-wInplaceReLU Best ISCIFAR10w564×15000003800009.228611.7760
SAGAN-32x32-wInplaceReLU Best FIDCIFAR10w564×15000004600009.206110.7781
SAGAN-128x128-woInplaceReLU Best ISImageNetw/o164×4100000098000031.593836.7712
SAGAN-128x128-woInplaceReLU Best FIDImageNetw/o164×4100000095000028.493634.7838
SAGAN-128x128-BigGAN Schedule Best ISImageNetw/o132×8100000082600069.535012.8295
SAGAN-128x128-BigGAN Schedule Best FIDImageNetw/o132×8100000082600069.535012.8295

从上表可以看到,"BigGAN Schedule" 变体在 ImageNet 上取得了显著更优的结果(IS 69.5350 / FID 12.8295),说明学习率解耦、EMA、谱归一化 epsilon 调整与无增强训练等设置对 ImageNet 这种大规模数据集的训练稳定性至关重要。

5.3 从 PyTorch-StudioGAN 转换的预训练模型

mmagic 还提供了从 PyTorch-StudioGAN 与 sagan_128_cvt_studioGAN.py。

模型数据集Inplace ReLUn_disc总迭代数IS(mmagic 评估)FID(mmagic 评估)IS(StudioGAN)FID(StudioGAN)
SAGAN-32x32 StudioGANCIFAR10w51000009.11610.20118.68014.009
SAGAN-128x128 StudioGANImageNetw1100000027.36740.116229.84834.726

表中Our Pipeline表示使用 mmagic 评估流程得到的结果,StudioGAN表示 PyTorch-StudioGAN 官方发布的结果。两套数值存在差异,原因在于评估细节的不同(见下一节)。

六、IS 与 FID 评估细节与差异说明

原文档明确指出,mmagic 的 IS 评估与 PyTorch-StudioGAN 存在两处实现差异:

  1. 特征提取器:使用 Tero 的 Inception(script module)进行特征提取;
  2. 图像缩放:在送入 Inception 之前,使用PIL 后端的 bicubic 插值进行缩放。

对于 FID 评估,mmagic 遵循BigGAN 的 pipeline——使用整个训练集提取 Inception 统计量;而 PyTorch-StudioGAN 仅使用随机选择的 50000 个样本。此外 mmagic 同样使用 Tero 的 Inception 进行特征提取。

6.1 下载预提取的 Inception 状态

为方便用户,mmagic 提供预提取的 inception 状态文件(CIFAR10 与 ImageNet1k 各一份)。用户也可以自行用以下命令提取这些状态(命令来自原文档,注意工具路径以仓库实际布局为准):

# 对于 CIFAR10 python tools/utils/inception_stat.py --data-cfg configs/_base_/datasets/cifar10_inception_stat.py --pklname cifar10.pkl --no-shuffle --inception-style stylegan --num-samples -1 --subset train # 对于 ImageNet1k python tools/utils/inception_stat.py --data-cfg configs/_base_/datasets/imagenet_128x128_inception_stat.py --pklname imagenet.pkl --no-shuffle --inception-style stylegan --num-samples -1 --subset train

另外,mmagic 的指标文档(docs/en/user_guides/metrics.md)还补充说明:FID 计算时真实特征会在测试时自动提取并保存在本地(默认缓存于MMAGIC_CACHE_DIR,即~/.cache/openmmlab/mmagic/),后续测试会自动读取缓存;参数变化会通过 hash 值标记特征文件。迁移到新机器时,可以直接复制缓存目录中的 pkl 文件并设置inception_pkl字段。

七、训练与推理实操

7.1 启动训练

mmagic 的训练入口为 tools/train.py,使用方式为:

python tools/train.py ${CONFIG_FILE}

例如训练 CIFAR10 32×32 的 SAGAN:

python tools/train.py configs/sagan/sagan_woReLUinplace_lr2e-4-ndisc5-1xb64_cifar10-32x32.py

训练过程中VisualizationHook会周期性(默认每 5000 次迭代)将固定噪声输入下的生成图像保存下来,便于直观观察生成质量的演进;checkpoint钩子会分别按 FID 最小、IS 最大保存最优权重。多卡训练可参考 tools/dist_train.sh 等分布式训练脚本。

7.2 测试与推理

使用 tools/test.py 即可基于训练配置与 checkpoint 进行评估:

python tools/test.py ${CONFIG_FILE} ${CHECKPOINT_FILE}

评估时配置中的metrics(IS-50k 与 FID-Full-50k)会被自动执行。对于 ImageNet 128×128 与 BigGAN Schedule 变体,评估默认使用 EMA 模型(sample_model='ema');CIFAR10 配置则使用原始模型(sample_model='orig')。

7.3 使用转换权重

如需直接使用 StudioGAN 转换模型,将 sagan_cvt-studioGAN_cifar10-32x32.py(或 sagan_128_cvt_studioGAN.py)作为配置,并指定从 configs/sagan/metafile.yml 对应条目中获取的权重路径即可加载预训练模型。

八、如何用本文配置进行二次开发

从源码结构可以推断,在 mmagic 中调整 SAGAN 主要有以下入口:

  1. 更换输出分辨率:修改output_scale/input_scale,并确认channels_cfg中存在对应分辨率的通道倍率(内置支持 32/64/128),否则需自定义channels_cfg;
  2. 调整注意力插入位置:修改attention_after_nth_block(支持 int 或 int 列表,传入小于 1 的索引会被忽略),观察自注意力对不同分辨率层级的影响;
  3. 调整谱归一化强度:通过with_spectral_norm、with_embedding_spectral_norm、sn_eps、sn_style组合控制;
  4. 切换初始化风格:init_cfg支持studio、BigGAN、SAGAN、SNGAN等类型;
  5. 启用/关闭 EMA:通过ema_config配置(如 BigGAN Schedule 变体的interval=1, momentum=0.999, start_iter=2000)。

九、引用

若在研究中使用了 SAGAN 或本仓库实现,可按原文档提供的信息引用:

@inproceedings{zhang2019self, title={Self-attention generative adversarial networks}, author={Zhang, Han and Goodfellow, Ian and Metaxas, Dimitris and Odena, Augustus}, booktitle={International conference on machine learning}, pages={7354--7363}, year={2019}, organization={PMLR}, url={https://proceedings.mlr.press/v97/zhang19d.html}, }

小结

本文以 configs/sagan/README.md 为主体,结合 mmagic 中 SAGAN 的源码实现、基础配置与指标文档,完整覆盖了 SAGAN 的算法思想、SAGAN/SNGANGenerator/ProjDiscriminator的模块结构与损失函数、CIFAR10 与 ImageNet 两套官方配置的逐项参数、BigGAN Schedule 变体的改进点、迭代计数换算规则、全部官方实验结果、StudioGAN 权重转换方法以及 IS/FID 的评估细节。读者可以据此直接复现官方结果,或基于上述配置入口进行自定义扩展。

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

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

相关推荐

上一篇:Langfuse 前端实践:使用函数式 setState 更新规避闭包过期与回调重建
下一篇:Vibe-Trading 的 AKShare 数据源实战:从免 Key 行情接口到回测回退链的完整解析

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

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

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

立即咨询