
CodeFormer 三阶段训练完整指南从 VQGAN 到可调控人脸复原 Transformer【免费下载链接】CodeFormer[NeurIPS 2022] Towards Robust Blind Face Restoration with Codebook Lookup Transformer项目地址: https://gitcode.com/gh_mirrors/co/CodeFormerCodeFormerNeurIPS 2022以码本查找 Transformer为核心将盲人脸复原拆解为三阶段渐进式训练先训 VQGAN 构建离散码本空间再训练 Code Sequence Prediction Module 完成低质量人脸到码序列的预测w0最后训练 Controllable Module 引入保真度权重 w 实现复原效果的可控调节。本文以仓库文档 docs/train.md 为主体骨架结合 options 目录下的三个训练配置与 basicsr 内的模型实现源码完整讲解数据集准备、三个阶段的具体训练命令、配置文件逐项解析以及预计算 latent code 加速训练的关键技巧帮助你从零复现 CodeFormer 的完整训练管线。CodeFormer 整体网络结构VQGAN 编码器/量化器/生成器与 Transformer 码序列预测、可控融合模块的协作关系图片来自 assets/network.jpg对应训练文档中的三阶段目标模块一、训练总览与前置条件1.1 三阶段训练管线根据 docs/train.mdCodeFormer 的训练分为三个递进阶段阶段训练内容配置文件名模型类型端口Stage IVQGAN离散码本自编码器options/VQGAN_512_ds32_nearest_stage1.ymlVQGANModel4321Stage IICodeFormerw0码序列预测模块options/CodeFormer_stage2.ymlCodeFormerIdxModel4322Stage IIICodeFormerw1可控模块options/CodeFormer_stage3.ymlCodeFormerJointModel4323三个阶段的命令统一通过 BasicSR 框架的训练入口basicsr/train.py启动例如 Stage I 的命令为python -m torch.distributed.launch --nproc_per_nodegpu_num --master_port4321 basicsr/train.py -opt options/VQGAN_512_ds32_nearest_stage1.yml --launcher pytorch其中--nproc_per_nodegpu_num需要替换为实际 GPU 数量配置文件中num_gpu: 8为官方训练设定可按硬件调整--master_port用于避免多任务并行时的端口冲突。1.2 PyTorch 版本注意事项文档给出了明确的版本适配提示For PyTorch versions 1.10, please replacepython -m torch.distributed.launchin the commands below withtorchrun.即PyTorch ≥ 1.10 时请将命令中的python -m torch.distributed.launch替换为torchrun其余参数保持不变。torch.distributed.launch在较新版本中已被torchrun取代这是多卡分布式训练时的首要适配点。1.3 框架基础文档明确指出项目基于 BasicSR 框架构建basicsr目录即内嵌的 BasicSR 代码训练、断点恢复resume、日志等机制均由 BasicSR 提供。训练配置通过-opt参数指定配置文件采用 YAML 格式由 basicsr/utils/options.py 解析。因此掌握后续各配置文件的字段含义是理解整个训练流程的关键。二、数据集准备FFHQ训练数据集为 FFHQFlickr-Faces-HQ需将其 512×512 版本存放至datasets/ffhq/ffhq_512目录与三个配置文件中dataroot_gt: datasets/ffhq/ffhq_512的默认路径保持一致。数据集在运行时由FFHQBlindDataset加载见 basicsr/data/ffhq_blind_dataset.py其核心职责是读取 GT 高清人脸图dataroot_gt指向的目录支持disk与lmdb两种io_backend对 GT 施加随机盲退化生成低质量输入in模糊、下采样、噪声、JPEG 压缩甚至随机灰度化、颜色抖动使模型学会处理未知退化在 Stage II/III 中若配置了latent_gt_path还会同步加载预计算的 GT 码序列latent_gt作为码序列预测的监督信号。Stage I 配置中use_corrupt: false表示 VQGAN 阶段不使用退化增强因为该阶段只需学习高清人脸的离散重建而 Stage II/III 的use_corrupt: true则会启用盲退化合成。三、Stage I训练 VQGAN3.1 启动命令python -m torch.distributed.launch --nproc_per_nodegpu_num --master_port4321 basicsr/train.py -opt options/VQGAN_512_ds32_nearest_stage1.yml --launcher pytorch3.2 配置文件逐项解析options/VQGAN_512_ds32_nearest_stage1.yml 的关键字段如下配置块字段值含义通用model_typeVQGANModel对应 basicsr/models/vqgan_model.py 中的训练模型通用num_gpu8官方使用 8 卡训练数据dataroot_gtdatasets/ffhq/ffhq_512高清 GT 路径数据in_size/gt_size512/512输入与 GT 分辨率数据mean/std[0.5,0.5,0.5]归一化参数映射到 [-1,1]数据use_corruptfalseVQGAN 阶段不做退化增强数据batch_size_per_gpu4单卡 batch8 卡总 batch32网络network_g.typeVQAutoEncoder见 basicsr/archs/vqgan_arch.py网络img_size/nf512/64输入尺寸 / 基础通道数网络ch_mult[1,2,2,4,4,8]各分辨率层通道倍增网络quantizernearest最近邻向量量化非 Gumbel网络codebook_size1024码本大小1024 个嵌入向量优化optim_g/optim_dAdam, lr7e-5, betas[0.9,0.99]生成器与判别器均用 Adam调度scheduler.typeCosineAnnealingRestartLR余弦退火重启调度total_iter1600000总迭代数 160 万训练ema_decay0.995指数移动平均保存params_ema权重训练net_d_start_iter30001判别器从第 30001 步才开始参与损失pixel_optL1Loss, weight1.0像素重建损失损失perceptual_optLPIPSLoss, weight1.0LPIPS 感知损失损失gan_optGANLoss(hinge), weight1.0Hinge GAN 损失日志save_checkpoint_freq1e4每 1 万步存一次 checkpoint配置中有一处注释值得注意# base_lr(4.5e-6)*bach_size(4)说明学习率是按 batch 缩放设计的eta_min: 6e-5的注释# no lr reduce in official vqgan code则表明余弦退火的最终学习率刻意不降太低以贴合原版 VQGAN 的行为。3.3 源码层原理VQAutoEncoder 与 VQGANModel从 basicsr/archs/vqgan_arch.py 可以看到VQAutoEncoder由三部分串联forward定义于第 385-389 行x self.encoder(x) # Encoder多层 ResBlock AttnBlock Downsample下采样 32 倍 quant, codebook_loss, quant_stats self.quantize(x) # VectorQuantizer最近邻量化到码本 x self.generator(quant) # Generator上采样重建 512×512 图像配置名VQGAN_512_ds32_nearest中的ds32即指 512→16 的 32 倍下采样量化后得到16×16256 个码 token这正是后续生成latent_gt_code1024.pth时size_latent 16的由来VectorQuantizer第 24-84 行通过欧氏距离找到最近码本向量并计算 commitment lossbeta0.25同时返回min_encoding_indices供 Stage II/III 作为 GT 码序列监督判别器VQGANDiscriminator是 PatchGAN 结构。训练模型VQGANModelbasicsr/models/vqgan_model.py在optimize_parameters中依次计算像素损失L1与感知损失LPIPS若迭代数超过net_d_start_iter再叠加hinge GAN 损失并通过calculate_adaptive_weight依据重建损失与 GAN 损失在最后一层上的梯度范数比自适应调节 GAN 权重上限 1.0再乘disc_weight0.8的 taming 系数码本损失l_codebook最终总损失为三部分之和。其中 EMA 机制ema_decay0.995会维护一份滑动平均权重net_g_ema并在保存时以params_ema键存储——后续 Stage II/III 加载 VQGAN 或generate_latent_gt.py读取权重时使用的正是params_ema。四、预计算训练集的 Latent CodeVQGAN 训练完成后文档建议预先为训练集计算码序列code sequence以显著加速后续两个阶段的训练——这样每次迭代都不必再前向推理 VQGAN 的编码器与量化器。4.1 运行命令python scripts/generate_latent_gt.py该脚本scripts/generate_latent_gt.py支持四个命令行参数参数默认值含义-i, --test_pathdatasets/ffhq/ffhq_512待编码的训练集 GT 目录-o, --save_root./experiments/pretrained_models/vqgan输出目录--codebook_size1024码本大小--ckpt_path./experiments/pretrained_models/vqgan/net_g.pthVQGAN 权重路径4.2 脚本原理脚本用ARCH_REGISTRY.get(VQAutoEncoder)构建与配置完全一致的 VQGAN512、nf64、ch_mult[1,2,2,4,4,8]、nearest、codebook_size1024加载checkpoint[params_ema]后对每张图做前向x, feat_dict vqgan.encoder(img, True) x, _, log vqgan.quantize(x) min_encoding_indices log[min_encoding_indices] # 16×16 的码索引每张图会生成原始图orig与水平翻转图hflip两组 16×16 码索引最终以字典形式保存为latent_gt_code1024.pthtorch.save于第 65-66 行。该文件随后通过配置中的latent_gt_path字段供 Stage II/III 读取数据集侧basicsr/data/ffhq_blind_dataset.py 第 49-53 行加载后按文件名索引取出对应码序列作为监督。4.3 不训练 VQGAN 的替代方案文档同时指出若不需要训练自己的 VQGAN可以直接获取官方发布的预训练 VQGANvqgan_code1024.pth及其对应码序列latent_gt_code1024.pth两者均随 v0.1.0 Release 发布。仓库内 scripts/download_pretrained_models.py 与 scripts/download_pretrained_models_from_gdrive.py 可辅助下载预训练权重下载后的存放路径与配置中vqgan_path/latent_gt_path的默认值./experiments/pretrained_models/vqgan/...保持一致即可。五、Stage II训练 CodeFormerw05.1 启动命令python -m torch.distributed.launch --nproc_per_nodegpu_num --master_port4322 basicsr/train.py -opt options/CodeFormer_stage2.yml --launcher pytorch5.2 配置文件逐项解析options/CodeFormer_stage2.yml 的核心字段配置块字段值含义通用model_typeCodeFormerIdxModel码序列预测训练模型见 basicsr/models/codeformer_idx_model.py数据use_corrupttrue启用盲退化数据blur_kernel_size/blur_sigma41/[1, 15]大核高斯模糊数据kernel_list/kernel_prob[iso,aniso]/[0.5,0.5]各向同性/异性核各 50%数据downsample_range[4, 30]大退化下采样 4~30 倍数据noise_range[0, 20]高斯噪声等级数据jpeg_range[30, 80]JPEG 质量数据latent_gt_path~注释中给出示例路径预计算码序列路径置~则运行时用network_vqgan现场生成网络network_g.typeCodeFormer主网络见 basicsr/archs/codeformer_arch.py网络dim_embd/n_head/n_layers512/8/9Transformer 嵌入维度 / 注意力头数 / 层数网络connect_list[32,64,128,256]与生成器融合的多尺度特征分辨率网络fix_modules[quantize,generator]冻结VQGAN 的量化器与生成器只训练 Transformer网络vqgan_path./experiments/pretrained_models/vqgan/vqgan_code1024.pth预训练 VQGAN 权重网络network_vqganVQAutoEncoder 配置无预计算 latent 时用于在线生成 GT 码优化optim_gAdam, lr1e-4仅生成器优化器无判别器调度scheduler.typeMultiStepLR, milestones[400000,450000], gamma0.5里程碑衰减调度total_iter500000总迭代 50 万训练use_hq_feat_loss/feat_loss_weighttrue/1.0码本特征HQ feat回归损失训练cross_entropy_loss/entropy_loss_weighttrue/0.5码索引交叉熵损失训练fidelity_weight0保真度权重Stage II 恒为 05.3 源码层原理CodeFormerIdxModel 的两类监督配置中的latent_gt_path: ~与network_vqgan段对应了两种 GT 码序列获取方式这在 basicsr/models/codeformer_idx_model.py 的init_training_settings第 46-57 行中有明确分支若配置了latent_gt_path直接从文件读取预计算码序列generate_idx_gt False否则若有network_vqgan构建并冻结一个 VQGANhq_vqgan_fix在每个迭代中在线对 GT 编码量化得到码索引generate_idx_gt True。训练时optimize_parameters第 86-124 行模型以w0, code_onlyTrue前向 CodeFormer只取码预测 logits 与低质量特征lq_feat总损失为HQ 特征回归损失mean((quant_feat_gt.detach() - lq_feat)^2) * feat_loss_weight迫使 Transformer 预测出的特征逼近码本特征通过get_codebook_feat由 GT 码索引取回交叉熵损失F.cross_entropy(logits.permute(0,2,1), idx_gt) * entropy_loss_weight直接监督 256 个 token 位置的码索引分类。由于fix_modules: [quantize,generator]VQGAN 的量化器与生成器参数被冻结basicsr/archs/codeformer_arch.py 第 172-176 行对fix_modules中的模块设置requires_gradFalse此阶段只优化 Transformer 码预测分支因此配置中也没有判别器与 GAN 损失。六、Stage III训练 CodeFormerw16.1 启动命令python -m torch.distributed.launch --nproc_per_nodegpu_num --master_port4323 basicsr/train.py -opt options/CodeFormer_stage3.yml --launcher pytorch6.2 配置文件逐项解析options/CodeFormer_stage3.yml 在 Stage II 基础上引入了关键变化配置块字段值与 Stage II 的差异通用model_typeCodeFormerJointModel联合训练模型见 basicsr/models/codeformer_joint_model.py数据blur_sigma[0.1, 10]小退化Stage II 为 [1,15]数据downsample_range[1, 12]小退化下采样Stage II 为 [4,30]数据noise_range[0, 15]小退化噪声数据jpeg_range[60, 100]小退化 JPEG数据*_large系列blur_sigma_large[1,15] 等大退化参数沿用 Stage II网络network_dVQGANDiscriminator, n_layers4新增判别器路径pretrain_network_g./experiments/pretrained_models/CodeFormer_stage2/net_g_latest.pth加载 Stage II 训练好的生成器路径pretrain_network_d./experiments/pretrained_models/CodeFormer_stage2/net_d_latest.pth加载 Stage II 判别器若有优化optim_g/optim_dAdam, lr5e-5双优化器学习率降为 5e-5调度scheduler.typeCosineAnnealingRestartLR, periods[150000], eta_min2e-5余弦退火调度total_iter150000总迭代 15 万训练scale_adaptive_gan_weight0.1缩放自适应 GAN 权重训练ema_decay0.997EMA 衰减略高于前两阶段损失pixel_opt/perceptual_opt/gan_optL1 / LPIPS / hinge GAN与 Stage I 相同的图像级损失Stage II 配置中fix_modules: [quantize,generator]未出现于 Stage III说明本阶段同时训练 Transformer 码预测分支与可控融合模块。数据侧的*_large系列参数第 34-38 行与FFHQBlindJointDatasetbasicsr/data/ffhq_blind_joint_dataset.py配合让每个 batch 同时包含小退化与大退化样本供训练策略按迭代阶段切换。6.3 源码层原理CodeFormerJointModel 与 w 参数调度basicsr/models/codeformer_joint_model.py 的optimize_parameters第 139-253 行实现了文档所述可控模块训练的核心策略——按迭代数调度小/大退化与保真度权重 w迭代区间小退化采样间隔small_per_n权重w说明≤ 400001全部小退化1.0先学保真40001 ~ 800001全部小退化1.3提高可控强度80001 ~ 120000120000全部大退化0.0退化为纯码序列预测 12000015混合退化1.3小/大退化混合训练当走小退化分支时模型以ww前向生成图像并计算像素L1、感知LPIPS与 GAN 损失大退化分支code_onlyTrue只计算码预测相关损失不做图像级损失——这正是CodeFormerJointModel相比CodeFormerIdxModel多出的联合优化含义。GAN 权重同样使用calculate_adaptive_weight自适应计算并额外乘以scale_adaptive_gan_weight0.1控制整体强度判别器训练从net_d_start_iter5001开始。6.4 w 参数如何实现可调控复原文档将 Stage III 定义为 Training Controllable Module其可控制性在 basicsr/archs/codeformer_arch.py 的CodeFormer.forward(self, x, w0, ...)第 223 行起中体现低质量特征经由 Transformer 得到码预测与融合特征后生成器各分辨率的特征与编码器对应层特征通过fuse_convs_dict[f_size]Fuse_sft_block第 151 行以SFTSpatial Feature Transform方式融合而w 正是控制码本预测特征与原始低质量特征混合比例的门控系数w0Stage II 产物完全依赖 Transformer 从码本预测的离散先验复原能力强但保真度低w1Stage III 训练的默认权重融合低质量输入的细节信息兼顾先验与保真推理时可取w∈[0,1]连续调节这也是 inference_codeformer.py 提供-w参数默认 0.7的底层依据。经过三阶段训练最终得到发布形态的codeformer.pth文档说明其预训练权重随 v0.1.0 Release 发布可用于 inference_codeformer.py、inference_colorization.py 与 inference_inpainting.py 等推理脚本。七、训练工程要点与常见问题7.1 分布式训练与端口三个阶段的命令分别使用--master_port4321/4322/4323且配置文件dist_params.port也对应设为29411/29412/29413见各 yml 末尾多任务并行或与现有训练任务共存时注意避免端口冲突。find_unused_parameters: true已开启适配了冻结模块fix_modules导致的参数未使用情况。7.2 断点恢复与权重加载path.resume_state用于恢复完整训练状态优化器、调度器、迭代数三个配置中均置~从头训练path.pretrain_network_g/pretrain_network_d用于加载预训练权重Stage II 通过vqgan_path加载 VQGANStage III 通过pretrain_network_g加载 Stage II 的net_g_latest.pthparam_key_g: params_ema指定读取 EMA 权重strict_load_g: false允许部分权重不匹配如 Stage II 的生成器主干与 Stage III 新增模块strict_load_d: true则严格要求判别器完整匹配。7.3 显存与效率优化预计算 latent code是官方推荐的提速手段配置latent_gt_path后如./experiments/pretrained_models/VQGAN/latent_gt_code1024.pth每次迭代省去 VQGAN 前向显存与耗时双双下降Stage III 的batch_size_per_gpu降为 3、num_worker_per_gpu降为 1是因为联合训练包含判别器与多分辨率融合显存开销更大EMA 机制ema_decay0.995 / 0.997在保存时同时输出params与params_ema测试与部署推荐使用 EMA 权重。7.4 日志与可视化配置中logger.use_tb_logger: true启用 TensorBoard 日志wandb字段留空~表示未启用 Weights Biasesprint_freq: 100控制训练日志打印频率save_checkpoint_freqStage I/II 为 1e4Stage III 为 5e3控制 checkpoint 保存频率。验证集段val:在三个配置中均以val_freq: 5e10关闭注释# no validation若需验证可自行取消注释并配置PairedImageDataset与 PSNR 指标metrics.calculate_psnr。训练产出统一落在experiments/目录下含配置名子目录、日志、可视化图与 checkpoint完整训练细节含恢复、验证等可进一步参考 BasicSR 框架的官方文档仓库侧的中文版本说明见 docs/train_CN.md。【免费下载链接】CodeFormer[NeurIPS 2022] Towards Robust Blind Face Restoration with Codebook Lookup Transformer项目地址: https://gitcode.com/gh_mirrors/co/CodeFormer创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考