资讯动态

mmagic 中 SAGAN 条件生成对抗网络的模型配置、训练与评估实战指南

发布时间:2026/9/29 6:14:32 来源:尧图企业网站定制
媒体生成计算机视觉深度学习人工智能大模型【免费下载链接】mmagicOpenMMLab 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 中复现与使用 SAGANSelf-Attention Generative Adversarial NetworksICML2019。读完本文你将掌握 SAGAN 的算法动机、其在 mmagic 中的源码结构与关键超参数、CIFAR10 与 ImageNet 两套官方训练配置的逐项含义、迭代计数换算规则以及使用 Inception ScoreIS与 Fréchet Inception DistanceFID评估生成质量的完整流程并了解如何将 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 ScoreIS从 36.8 提升到 52.52并将 Fréchet Inception DistanceFID从 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注册名SAGANGeneratorsagan_discriminator.py投影判别器ProjDiscriminatorsagan_modules.py生成器 ResBlockSNGANGenResBlock、条件归一化SNConditionNorm等基础模块。从类注释与代码实现可以推断mmagic 的 SAGAN 实现同时融合了三个相关工作组件来源工作作用SAGANSelf-Attention GAN自注意力长期依赖建模SNGANGeneratorSpectral Normalization GANSNGAN生成器谱归一化ProjDiscriminatorcGANs with Projection DiscriminatorProj-GAN投影式条件判别器自注意力模块SelfAttentionBlock则复用了 BigGAN 实现中的模块见 biggan_modules.py 相关定义在基础配置中通过attention_cfgdict(typeSelfAttentionBlock)注入。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_lossloss_disc_fake relu(1 D(fake))的均值加上loss_disc_real relu(1 - D(real))的均值生成器损失gen_lossloss_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 生成器 SNGANGeneratorSNGANGenerator的关键设计是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(typeSNGANGenResBlock)用户可通过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_classes0时条件归一化层会自动退化为无条件版本。此外init_weights支持多种初始化风格STUDIOPytorch-StudioGAN正交初始化、BIGGANxavier_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_normsn_styletorch谱归一化实现风格torchPyTorch 官方实现或ajbrockBigGAN-PyTorch 实现sn_eps1e-12谱归一化操作的 epsilonnorm_eps1e-4条件/非条件归一化层的 epsilonauto_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×32CIFAR10基础配置为例model dict( typeSAGAN, data_preprocessordict(typeDataPreprocessor), num_classes10, generatordict( typeSNGANGenerator, num_classes10, output_scale32, base_channels256, attention_cfgdict(typeSelfAttentionBlock), attention_after_nth_block2, with_spectral_normTrue), discriminatordict( typeProjDiscriminator, num_classes10, input_scale32, base_channels128, attention_cfgdict(typeSelfAttentionBlock), attention_after_nth_block1, with_spectral_normTrue), generator_steps1, discriminator_steps5)128×128ImageNet基础配置与之对应num_classes1000、生成器output_scale128、base_channels64、attention_after_nth_block4判别器input_scale128、base_channels64、attention_after_nth_block1且generator_steps1、discriminator_steps1。可见在 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(typestudio)采用 Pytorch-StudioGAN 风格的正交初始化data_preprocessordict(output_channel_orderBGR)CIFAR 图像为 RGB需转换为 BGR 通道顺序train_cfg dict(max_iters100000 * disc_step)总迭代数 100000 × 5train_dataloader dict(batch_size64)单卡 batch size 64优化器生成器与判别器均使用 Adamlr0.0002betas(0.5, 0.999)VisualizationHook每 5000 次迭代可视化一次固定输入的生成结果fixed_inputTruevis_kwargs_listdict(typeGAN, namefake_img)。3.3 ImageNet 128×128 训练配置配置文件 sagan_woReLUinplace_Glr1e-4_Dlr4e-4_ndisc1-4xb64_imagenet1k-128x128.py 的关键设置生成器与判别器学习率解耦生成器Adam lr0.0001、判别器Adam lr0.0004两者betas(0.0, 0.999)discriminator_steps1ndisc1判别器与生成器交替更新train_cfg dict(max_iters1000000, val_interval10000, dynamic_intervals[(800000, 4000)])总迭代 100 万次每 1 万次迭代验证一次80 万次迭代后验证间隔动态调整为 4000train_dataloader dict(batch_size64)注释标注为 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 点差异谱归一化使用eps1e-8不使用 SyncBNauto_sync_bnFalse条件归一化cBN中的 embedding 层不使用谱归一化with_embedding_spectral_normFalse在特定迭代开始启用 EMAema_configdict(interval1, momentum0.999, start_iter2000)权重初始化使用xavier_uniforminit_cfg dict(typeBigGAN)不进行数据增强继承imagenet_noaug_128.py数据集。此外该配置的生成器还设置了norm_eps1e-5、sn_eps1e-8判别器sn_eps1e-8可视化 Hook 同时输出 EMA 与原始权重模型的生成结果sample_modelema/origtarget_keys[ema.fake_img, orig.fake_img]评估指标则使用 EMA 模型sample_modelema。四、评估指标配置所有训练配置末尾都挂载了 IS 与 FID 两个指标例如inception_pkl ./work_dirs/inception_pkl/cifar10-full.pkl metrics [ dict( typeInceptionScore, prefixIS-50k, fake_nums50000, inception_styleStyleGAN, sample_modelorig), dict( typeFrechetInceptionDistance, prefixFID-Full-50k, fake_nums50000, inception_styleStyleGAN, inception_pklinception_pkl, sample_modelorig) ] default_hooks dict( checkpointdict( save_best[FID-Full-50k/fid, IS-50k/is], rule[less, greater]))要点说明fake_nums50000评估时生成 5 万张假图inception_styleStyleGAN使用 Tero 的 Inception V3 script module 提取特征详见 docs/en/user_guides/metrics.mdinception_pklFID 需要真实数据集的 Inception 特征统计量预先保存为 pkl 可避免每次评估重复提取default_hooks.checkpoint同时保存 FID 最小与 IS 最大的两个最优 checkpointrule[less, greater]。关于 Inception V3 与图像缩放方式的选择mmagic 的指标文档明确指出这两者会显著影响最终 IS 分数因此强烈推荐使用 Tero 的 script model加载需要torch 1.6并采用Pillow 后端的 Bicubic 插值进行缩放。对应配置中可通过resize_method与use_pillow_resize设置缩放方式通过inception_style选择StyleGANTero 模型或PyTorchtorchvision 实现在无网络环境下可下载 Inception 权重并通过inception_path指定。五、迭代计数规则与实验结果5.1 迭代计数换算原文档特别强调mmagic 实现的迭代计数规则与其他代码库不同。若需与其他代码库对齐可使用如下换算公式total_iters (biggan/pytorch studio gan) our_total_iters / dist_step其中dist_step即配置中的disc_step判别器每轮更新次数。例如 CIFAR10 配置中disc_step5、max_iters500000对应其他代码库的 100000 次迭代。5.2 官方训练结果以下是 mmagic 官方在 CIFAR10 与 ImageNet 上训练的 SAGAN 模型结果模型权重可通过 configs/sagan/metafile.yml 中对应条目的Weights字段获取模型数据集Inplace ReLUdist_step总 batch size总迭代数*最佳迭代ISFIDSAGAN-32x32-woInplaceReLU Best ISCIFAR10w/o564×15000004000009.321710.5030SAGAN-32x32-woInplaceReLU Best FIDCIFAR10w/o564×15000004800009.31749.4252SAGAN-32x32-wInplaceReLU Best ISCIFAR10w564×15000003800009.228611.7760SAGAN-32x32-wInplaceReLU Best FIDCIFAR10w564×15000004600009.206110.7781SAGAN-128x128-woInplaceReLU Best ISImageNetw/o164×4100000098000031.593836.7712SAGAN-128x128-woInplaceReLU Best FIDImageNetw/o164×4100000095000028.493634.7838SAGAN-128x128-BigGAN Schedule Best ISImageNetw/o132×8100000082600069.535012.8295SAGAN-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总迭代数ISmmagic 评估FIDmmagic 评估ISStudioGANFIDStudioGANSAGAN-32x32 StudioGANCIFAR10w51000009.11610.20118.68014.009SAGAN-128x128 StudioGANImageNetw1100000027.36740.116229.84834.726表中Our Pipeline表示使用 mmagic 评估流程得到的结果StudioGAN表示 PyTorch-StudioGAN 官方发布的结果。两套数值存在差异原因在于评估细节的不同见下一节。六、IS 与 FID 评估细节与差异说明原文档明确指出mmagic 的 IS 评估与 PyTorch-StudioGAN 存在两处实现差异特征提取器使用 Tero 的 Inceptionscript module进行特征提取图像缩放在送入 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 的 SAGANpython 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}评估时配置中的metricsIS-50k 与 FID-Full-50k会被自动执行。对于 ImageNet 128×128 与 BigGAN Schedule 变体评估默认使用 EMA 模型sample_modelemaCIFAR10 配置则使用原始模型sample_modelorig。7.3 使用转换权重如需直接使用 StudioGAN 转换模型将 sagan_cvt-studioGAN_cifar10-32x32.py或 sagan_128_cvt_studioGAN.py作为配置并指定从 configs/sagan/metafile.yml 对应条目中获取的权重路径即可加载预训练模型。八、如何用本文配置进行二次开发从源码结构可以推断在 mmagic 中调整 SAGAN 主要有以下入口更换输出分辨率修改output_scale/input_scale并确认channels_cfg中存在对应分辨率的通道倍率内置支持 32/64/128否则需自定义channels_cfg调整注意力插入位置修改attention_after_nth_block支持 int 或 int 列表传入小于 1 的索引会被忽略观察自注意力对不同分辨率层级的影响调整谱归一化强度通过with_spectral_norm、with_embedding_spectral_norm、sn_eps、sn_style组合控制切换初始化风格init_cfg支持studio、BigGAN、SAGAN、SNGAN等类型启用/关闭 EMA通过ema_config配置如 BigGAN Schedule 变体的interval1, momentum0.999, start_iter2000。九、引用若在研究中使用了 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 的评估细节。读者可以据此直接复现官方结果或基于上述配置入口进行自定义扩展。赞分享媒体生成计算机视觉深度学习人工智能大模型【免费下载链接】mmagicOpenMMLab 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点击查看免费下载相关推荐containerd 运维实战指南从 systemd 托管到插件级配置的完整手册containerd 运维实战指南从 systemd 托管到插件级配置的完整手册 导读 本指南面向运维与管理员Ops and Admins系统讲解 co媒体生成计算机视觉深度学习人工智能大模型4 步让 Dify 工作流接上外部 APIHTTP 请求节点实战4 步让 Dify 工作流接上外部 APIHTTP 请求节点实战 Awesome Dify Workflow 是一个 Dify 工作流 DSL 合集本文基于媒体生成计算机视觉深度学习人工智能大模型MMagic 中的 BigGAN 条件图像生成原理、配置、训练与采样实战指南MMagic 中的 BigGAN 条件图像生成原理、配置、训练与采样实战指南 导读 本文以 configs/biggan/README.md https://媒体生成计算机视觉深度学习人工智能大模型上一篇Langfuse 前端实践使用函数式 setState 更新规避闭包过期与回调重建下一篇Vibe-Trading 的 AKShare 数据源实战从免 Key 行情接口到回测回退链的完整解析创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑