首页
/ Faceswap Train 插件包深度解析:Model 与 Trainer 的基类架构、状态管理与配置体系

Faceswap Train 插件包深度解析:Model 与 Trainer 的基类架构、状态管理与配置体系

2026-09-06 15:28:47作者:裘晴惠Vivianne

本文以 Faceswap 仓库的 train 插件包文档 为核心,完整拆解 Faceswap 训练体系的两大支柱:plugins/train/model(模型插件基类及其 inference / io / model / settings / state / update 六大基础模块)与 plugins/train/trainer(训练循环基类、原版与分布式训练器、训练器配置)。读完后,你将理解 Faceswap 模型插件的继承约定、.keras 模型文件与 state 状态文件的读写机制、优化器状态注入、混合精度切换、Keras 2 遗留模型迁移,以及训练配置(损失函数、优化器、掩码、数据加载、数据增强)每一项参数的实际取值与默认值。

1. Train 包总体结构

docs/full/plugins/train.rst 定义了 train 包的两层结构:

从源码结构看,plugins/train/model/ 目录下还存在大量具体模型插件(dfaker.pydlight.pyunbalanced.pyvillain.pyphaze_a.pyrealface.pyiae.pylightweight.pydfl_h128.pydfl_sae.py 等),每个插件都配套一个 *_defaults.py 文件存放该模型专属的默认配置,例如 dfaker_defaults.py。这些插件全部继承下文要介绍的 ModelBase

训练入口由 faceswap.py 统一分发,其中 TrainArgs 对应 train 子命令("Train a model for the two faces A and B"),即命令行入口 faceswap.py train ...

2. Model 基类:ModelBase 的职责与生命周期

model.py 中的 ModelBase所有模型插件必须继承的基类。其 __init__ 接收三个参数:

  • model_dir:模型保存位置的完整路径;
  • arguments:由 Faceswap 命令行参数生成的参数命名空间;
  • predictTrue 表示以推理(convert)方式加载模型,False 表示以训练方式加载,默认 False

初始化时基类完成了若干关键动作(见 model.py#L45-L87):

  1. 设置 input_shape 占位(模型插件必须在自己的 __init__ 中覆写为 3 维 shape 元组;若 A/B 两侧输入尺寸不同,则应赋值为两个 shape 元组的列表);
  2. 通过 cfg.load_config(config_file=...) 加载 train_config.py 定义的全局训练配置;
  3. 前置校验:若选择了 Penalized Mask LossLearn Mask 却未选择任何 Mask(mask_typenone),直接抛出 FaceswapError 终止;
  4. 实例化 IO(模型存取)、State(状态文件)、Settings(精度与后端设置)三大组件。

2.1 关键属性约定

属性 含义
model 编译后的 Keras 模型本体
name 插件名,直接取自插件文件名(去掉扩展名并转小写),即模型在磁盘上的标识
model_name Keras 模型名,默认同 name;含多套架构的插件可覆写
input_shapes / output_shapes 模型全部输入/输出张量的形状列表
iterations 模型已训练的总迭代次数(来自 state 文件)
freeze_layers / load_layers 可冻结/可加载权重的层名,默认均为 ["encoder"]
coverage_ratio 训练裁剪覆盖率,取值为 cfg.coverage() / 100,注释中明确了像素级换算公式:(original_size * coverage_ratio // 2) * 2,以保证偶数像素尺寸
color_order 图像通道顺序,默认 bgr

插件侧必须实现 build_model(inputs)NotImplementedError 占位),接收 _get_inputs() 返回的 face_in_a / face_in_b 两个 keras.layers.Input 张量,构建 A/B 双侧(双输入)自编码器架构。

2.2 build():构建或加载模型

build() 的执行路径是:

  • io.model_exists 为真:从磁盘 load() 已有 .keras 文件;若处于推理模式,则交给 Inference 类裁剪成单侧推理模型;
  • 否则:校验 input_shape,生成 A/B 两个输入张量,构建全新模型。若未启用混合精度且非 --summary 模式,会调用 Settings.get_mixed_precision_layers() 记录可切换混合精度的层名并写入 state;
  • 训练模式下调用 _compile_model()——按源码注释,这已不是传统 Keras compile,而是「加载并冻结权重」:检查混合精度是否需要重建(state.model_needs_rebuild)、实例化 Weights 处理 --load-weights--freeze-weights、然后冻结;
  • 最后输出模型 summary(--summary 时直接打印到 stdout,否则写入 verbose 日志,包含所有子模型)。

此外 _check_multiple_models() 会检查模型目录:目录中若已存在其他插件的模型文件、或同时存在多个插件类型的模型文件,都会抛出 FaceswapError 要求更换目录——这是 IO.multiple_models_in_folder 属性按 os.path.commonprefix 判定的结果。

3. inference 模块:训练模型到推理模型的裁剪

inference.py 中的 Inference 类解决一个具体问题:Faceswap 训练模型是 A/B 双输入的对称网络,但换脸时只需要其中一个方向的半边网络。Inference(saved_model, switch_sides) 的参数:

  • saved_model:已保存的完整训练模型;
  • switch_sidesTrue 表示 B→A 方向换脸,False 表示 A→B 方向。

其内部流程(可对照 inference.py#L29-L205):

  1. _get_input:断言模型输入数为 2,取 side_idx(B→A 为 0,A→B 为 1)对应的输入张量;
  2. _get_valid_layer_inputs:对每个非 InputLayer 层,通过 Keras 内部 _inbound_nodes 找出该层在所选换脸方向下的有效输入层列表;
  3. _get_output_layer:从 model.output 的 Keras history 中找出输出层——支持两种结构:共享解码器(全部输出指向同一层)与拆分解码器(A/B 各自独立输出层,按 side_idx 取前半或后半);
  4. _backwards_recurse:从输出层反向递归收集属于该换脸路径的全部层(带 seen 集合防环);
  5. __call__:按拓扑依赖逐层调用已收集到的层实例,重建出一个名为 {原模型名}_inference 的单输入 Keras 推理模型返回。

这意味着 convert 阶段不需要加载整个双塔网络的全部权重路径,只裁剪出实际用到的半边,从而降低显存占用。

4. io 模块:模型存取、优化器注入与备份

io.pyIOWeightsOptimizerMigrate 三个类与 get_all_sub_models() 工具函数组成。

4.1 IO:文件布局与保存策略

  • 文件名约定{model_dir}/{插件名}.keras(如 original.keras),model_exists 据此判断;
  • 加载容错load() 捕获 RuntimeError/KeyError 中 "unable to open object" 类错误时提示模型可能已损坏、建议使用 Restore Tool;捕获到 "parameter name can't contain '.'" 或缺少 Conv2D/DepthwiseConv2D 类时,自动调用 PatchKerasConfig 修补配置后重试加载(见 io.py#L153-L202);
  • 遗留迁移_update_legacy() 在目录中找不到 .keras 文件时查找同名 .h5,若有则交给 Legacy 类完成 Keras 2 → Keras 3 升级;
  • 优化器状态注入_save_optimizer() 把 PyTorch 优化器的 state_dict() 序列化为 optimizer.pt,以追加模式写入 .keras 文件的 zip 包中;load_optimizer() 反过来读取该条目,若只有旧的 config.json(Keras 优化器)则走 OptimizerMigrate.convert() 迁移。这正是 save_optimizer 配置项(never / always / exit)在保存阶段的落点;
  • 保存与备份save() 先执行 _save_model()(按 save_optimizer 策略决定是否写优化器),随后 _maybe_backup() 计算自上次保存以来的平均损失_get_save_average() 会重置 history),仅当该平均值低于历史最低(state.lowest_avg_loss)时才调用 Backup.backup_model() 备份 .keras 与 state 文件——源码注释明确说明这是一种防止把 NaN 损坏模型当备份保留的启发式手段,但并非绝对可靠(若损坏恰发生在接近保存点、或调整过损失权重,仍可能备份到坏模型);
  • 快照snapshot() 在模型保存后的下一次迭代触发(所以迭代数减 1),调用 Backup.snapshot_models()

4.2 Weights:冻结与权重迁移

Weights 处理两个 CLI 开关:

  • --freeze-weightsfreeze()统一解冻所有子模型(因为 layer.trainable 检查在冻结后仍可能误报 True),再把插件 freeze_layers(默认 ["encoder"])中实际存在的层设为 trainable = False;找不到对应层名时输出警告;
  • --load-weightsload() 要求权重文件必须是有效 .keras 路径(_check_weights_file 校验存在性与扩展名);仅在新建模型时生效——若模型已存在(恢复训练)则忽略权重文件并警告;加载时要求权重模型顶层 name 与当前模型一致(不允许跨插件加载),逐层比对 shape,shape 不符的层跳过并告警;若最终没有任何权重加载成功则抛错。

4.3 OptimizerMigrate:Keras 优化器到 Torch state_dict 的字段映射

这是恢复老模型训练连续性的关键实现。OptimizerMigrate 内置了一张 Keras 优化器到 Torch 优化器状态字段的映射表:

Keras 优化器 Keras 侧字段 Torch 侧字段
AdaBeliefOptimizer _momentums, exp_avg exp_avg, exp_avg_var
adam / adamw _momentums, exp_avg exp_avg, exp_avg_sq
adamax _m, _u exp_avg, exp_inf
lion _momentums exp_avg(无 step)
nadam _momentums, exp_avg, _u_product exp_avg, exp_avg_sq, mu_product
rmsprop _velocities square_avg

迁移输出为 {"version": 0.5, "optimizer": {"state": ..., "param_groups": ...}, "scaler": ...} 结构:param_groups 区分 decay / no_decay 两组参数(对应 weight decay 与 bias);若优化器是 LossScaleOptimizer 包装(混合精度),还会把 dynamic_scaledynamic_growth_stepsstep_counter 映射为 Torch 的 scale / growth_interval / _growth_tracker。若 Keras 内部结构变化导致字段缺失,会明确抛出 RuntimeError 提示 "Keras version may have changed internal structure"。

5. state 模块:状态文件与恢复训练的语义

state.py 中的 State 负责 <model_name>_state.json 状态文件。state 文件中保存:

  • name:插件名;
  • sessions:训练会话字典,每个会话记录 timestampno_logsbatchsizeiterationsconfig(仅包含可更新项);
  • lowest_avg_loss:历史最低保存间隔平均损失(用于 io 的备份判定);
  • iterations:总迭代数;
  • mixed_precision_layers:可混合精度切换的层名列表;
  • lr_finder:学习率查找器发现的值(-1 表示无);
  • config:完整训练配置快照。

5.1 fixed 与 updatable:配置项的双轨语义

恢复训练时,_update_config() 的核心逻辑是以 state 文件中的配置为权威

  • fixed=True 的配置项(如 centeringcoverageoptimizer):若用户当前配置与 state 不一致,用 state 中的值覆盖用户配置opt.set(old_val)),即改这些项对已有模型无效;
  • fixed=False 的配置项(如 learning_rateloss_functionmixed_precision):允许运行时更新,新值写回 state,并记入 _updatable_options;若变更项属于 rebuild_tasks(目前是 mixed_precision),置 model_needs_rebuild = True 触发模型重建;
  • 对 state 中缺失的新增配置项,优先采用 legacy_defaultscentering="legacy"coverage=62.5mask_loss_function="mse"optimizer="adam"mixed_precision=False)——源码注释解释了原因:这些新默认值可能对老模型有害,所以老模型按旧行为补齐。

5.2 遗留 state 配置迁移

_update_legacy_config() 实现了老 state 文件字段到新字段的自动改写,值得注意的映射包括(见 state.py#L200-L297):

旧字段 新处理
dssim_loss 转成 loss_function(true→ssim,false→mae),删除旧项
mask_type(缺 learn_mask 时) 补充 learn_mask,值为 mask_type 非 None
mask_typefacehull/dfl_full 替换为 components(这两种掩码已移除)
l2_reg_term 转为 loss_function_2=mse + loss_weight_2=旧值
clipnorm / autoclip 转为 gradient_clipping 类型
混合精度层名含 "." 点号替换为下划线(Keras 3.12 起点号层名会 KeyError)

模型插件专属默认值则从 plugins.train.model.{name}_defaults 模块动态导入(_get_model_options()),这解释了为什么每个模型插件都有配套 *_defaults.py

6. settings 模块:混合精度的实现细节

settings.pySettings 在模型启动前完成后端设置。混合精度的启用条件是 not is_predict and mixed_precision(推理不启用),核心方法是:

  • _set_keras_mixed_precision():通过 Keras 3 的 dtype_policies.DTypePolicy("mixed_float16" | "float32") + k_config.set_dtype_policy() 全局切换;
  • get_mixed_precision_layers():采用两次构建策略——先在 CPU 上以 mixed_float16 策略构建一次模型,从 get_config()["layers"] 递归收集 dtype 为 mixed_float16 的层名(Functional/Sequential 子模型会递归进入),然后清会话、切回 float32 再正式构建,返回 FP32 模型与兼容层名清单供 state 保存;
  • check_model_precision():在恢复训练且精度模式发生变化时,Keras 无法直接改 dtype,于是改 config 里各兼容层的 dtype 字符串 → from_config 重建新模型 → set_weights() 把旧权重搬过去;对从未记录过兼容层名清单的更老模型文件,则明确拒绝「FP32 → 混合精度」的切换并回退到全精度(只打印警告)。

loss_scale_optimizer() 则是把任意优化器包装进 optimizers.LossScaleOptimizer 的工厂方法,供混合精度下做动态损失缩放。

7. update 模块:Keras 2 与旧版 Keras 3 模型的自动升级

update.py 提供两个迁移类:

7.1 Legacy:Keras 2 .h5 → Keras 3 .keras

Legacy.__init__(model_path) 直接执行完整升级流程:

  1. _get_model_config():绕过 Keras 3 的加载器,用 h5py 直接从 .h5 的 attrs 中读取 keras_versionmodel_config,校验主版本必须为 2;
  2. _update_layers() 递归处理每层配置:
    • LeakyReLUalpha 参数改名为 negative_slope
    • 名为 visual 的层(CLiP 编码器中的 MultiHeadAttention)从同名 _state.json 中读取 enc_architecture/enc_scaling,通过 plugins.train.model.phaze_a._MODEL_MAPPING 找到 ViT 网络信息后重建新 config(缩放按 16 的倍数对齐);
    • TFOpLambda 层替换为自研 ScalarOp 层,仅支持 multiply/truediv/add/subtract 四种运算;
    • DepthwiseConv2D/SeparableConv2D/Conv2DTranspose 删除 Keras 3 不认的 groups 参数;
    • 错误的 dtype 存储统一改写为 DTypePolicy 结构;
    • 共享 Functional 模型的 inbound node index 减 1(修复 Keras 3 遗留加载 bug),嵌套 output_layers 展平(Keras 3 不接受嵌套输出);
  3. _archive_model():把旧模型目录整体改名归档为 {model_dir}_fs2_backup(目标目录已存在且非空则报错),再 _restore_files()*_state.json*_logs 目录复制回新目录。

7.2 PatchKerasConfig:旧版 Keras 3 模型文件的原地修补

针对 3.12 之前的旧 Keras 3 .keras 文件(zip 内 config.json),处理两类不兼容:

  • _update_nn_blocks():把 lib.model.nn_blocks 下自定义的 Conv2D/DepthwiseConv2D 模块引用改写为 keras.layers 版本(对应项目侧不再继承 Keras 层、改为内部调用底层的重构);
  • _update_dot_naming():层名/config 名/inbound keras_history 中的点号统一替换为下划线(主要影响使用 CLiP 编码器的模型)。

修补后以 ZIP_DEFLATED 压缩重写整个 zip。IO.load() 在捕获到相应异常时会自动触发本类并重试,用户无感知。

8. Trainer 包:训练循环、原版与分布式

8.1 TrainConfig 与 TrainerBase

base.py 定义了训练配置数据类 TrainConfig

字段 默认值 说明
folders 必填 输入文件夹列表,按处理顺序(如 [A, B]
batch_size 必填 每个 loader 每批加载的数据量
augment_color 必填 是否做颜色增强
flip 必填 是否做水平翻转
warp 必填 是否做 warp
cache_landmarks 必填 是否缓存对侧 landmarks 用于 warp-to-landmarks
lr_finder False 是否启用学习率查找器
snapshot_interval -1(禁用) 快照间隔迭代数

TrainerBase 是抽象基类:持有 modelbatch_sizeconfigsampler(由子类 get_sampler() 提供,类型限定为 RandomSamplerDistributedSampler),register_loss() 会把配置的 LossCollatoradd_module("loss_func", loss) 挂到 Keras 模型上——这正是分布式模块能直接调用 self._keras_model.loss_func(...) 的原因。抽象方法 train_batch(inputs, targets, optimizer, meta) 要求子类执行一次完整的前向+反向,返回按 A/B/… 顺序排列的各侧损失。

8.2 original:标准训练器

original.py 是最小实现:

  • get_sampler() 返回 torch.utils.data.RandomSampler
  • _forward()predictions = self.model.model(inputs, training=True),按 num_outputs = len(predictions) // num_sides 切片,逐侧调用 self.loss_func(targets_i, predictions_i, meta_i) 得到各侧损失;
  • _backwards_and_apply()sum(x.total for x in loss) 汇总各侧总损失后 optimizer.backward(total_loss) + optimizer.step()
  • train_batch() = 前向 + 反向。

8.3 distributed:DataParallel 多卡训练

distributed.pyTrainer 继承 original,核心改动:

  1. WrappedModel:把双输入 Keras 模型包成单入口 torch.nn.Module,其 forward(inputs, targets, meta_dict)每张卡上完成前向并计算各侧损失,返回损失字典列表(过滤 None 项);
  2. _validate_batch_size():批次小于 GPU 数时自动提升到 GPU 数并告警;批次不能被 GPU 数整除时提示更优值((batch // gpus) * gpus 或其加一个 GPU 数)——因为 DataParallel 按卡均分批次;
  3. _set_distributed():用 torch.nn.DataParallel(WrappedModel(...)) 包装模型(CUDA_VISIBLE_DEVICES 已由 -X 命令行参数预先设置),并拦截 Torch 关于 GPU 显存不均衡的 UserWarning、裁剪掉与 Faceswap 无关的后缀提示;
  4. _mean_loss():由于 DataParallel 对非标量返回按卡列表,递归对 tensor 取 mean(),再把合并结果还原为 BatchLoss 对象。

GPU 不均衡告警的处理逻辑在 _handle_torch_gpu_mismatch_warning() 中,截断到 "You can do so by" 之前的有效信息再写入日志。

9. 训练配置体系:train_config.py 全参数解析

train_config.py 是全局训练配置(对应 config.iniglobaltrainer.* 节区)。文档顶部有一句总注记:除特别说明外,此处修改的取值仅对新建模型生效——这正是 State 中 fixed=True 项的语义。load_config(config_file) 用全局标志 _IS_LOADED 保证只加载一次,且支持传入自定义 ini 文件。

9.1 人脸(face)组

配置项 类型/默认 范围/选项 fixed 说明
centering str,face face/head/legacy 训练图裁剪中心:face 按姿态校正的脸中心;head 按头中心(仅当最终换脸要包含头发、且掩码覆盖头发时选择);legacy 为原始提取技术,鼻尖附近无校正,脸缘可能出框
coverage float,100.0 62.5–100 训练使用对齐图像的裁剪比例。face 中心建议 >75%;head 中心建议 100%;legacy 中心下 62.5%≈眉到眉、75%≈鬓角到鬓角、87.5%≈耳到耳、100%≈全景照
vertical_offset int,0 -25–25 垂直位移百分比,负值露出更多下巴,正值露出更多额头

9.2 初始化(initialization)组

配置项 默认 说明
icnr_init False ICNR 平铺初始化,配合亚像素/pixel shuffler 缓解棋盘格效应
conv_aware_init False 卷积感知初始化,可缓解梯度消失/爆炸、加快收敛;注意会多占 VRAM(首跑可降批次)、不支持多卡(需单卡启动后再切多卡恢复)、建模耗时数分钟

9.3 学习率查找器组

配置项 默认 范围 fixed 说明
lr_finder_iterations 1000 100–10000(步长 100) 查找过程处理的迭代数
lr_finder_mode set set/graph_and_set/graph_and_exit 分别对应「用发现值直接训练」「输出图并训练」「输出图后退出」;已有模型恒按 set 处理
lr_finder_strength default default/aggressive/extreme 激进度:extreme 取最高最优学习率,梯度爆炸风险显著上升

查找到的学习率最终通过 State.add_lr_finder() 写入 state 文件的 lr_finder 字段。

9.4 网络(network)与 convert 组

配置项 默认 fixed 说明
reflect_padding False 卷积使用反射填充替代零填充,减少图像边缘伪影
mixed_precision False float16/float32 混合精度,主要惠及计算能力 7.0+(Tensor Core)的 NVIDIA 卡;老卡只有显存/带宽收益
nan_protection True 检测到 NaN 立即停止训练,最后一次保存不含 NaN,模型尚有机会抢救
convert_batchsize 16 convert 阶段 GPU 批大小(1–32),OOM 时调小

9.5 Loss 节区:损失函数与掩码

Loss 是 GlobalSection dataclass。损失函数候选共 13 种(_LOSS_HELP 字典逐项有论文级说明),其中 flip、三种 LPIPS 与 none 被标记为不可作主损失_NON_PRIMARY_LOSS):

  • 主损失loss_function,默认 ssim):候选为 fflgmsdl_inf_normlaplosslogcoshmaemsems_ssimpixel_gradient_diffsmooth_lossssim
  • 第二/三/四损失loss_function_2/3/4,默认 mse/none/none):可取全部候选;对应 loss_weight_2/3/4(默认 100/0/0,范围 0–400,按百分比缩放,0 即禁用);
  • mask_loss_function(默认 mse,可选 mae/mse):学习掩码时使用的损失;
  • eye_multiplier(默认 3,1–40)、mouth_multiplier(默认 2,1–40):眼部/嘴部相对主损失的优先级倍数,依赖 penalized_mask_loss 开启
  • penalized_mask_loss(默认 True):损失按掩码加权,掩码外区域的重建误差被忽略;
  • mask_type(默认 extended):可选 nonecomponentsextended(凸包外扩到额头),外加 PluginLoader.get_available_extractors("mask") 动态提供的 NN 掩码插件(文档说明中列出了 bisenet-fp_facebisenet-fp_headvgg-clearvgg-obstructedunet-dflcustom_facecustom_head 等);
  • mask_dilation(默认 0.0,-5–5):负值腐蚀、正值膨胀掩码,按掩码尺寸百分比计;
  • mask_blur_kernel(默认 3,0–9):对掩码做高斯模糊平滑边缘,偶数自动进位为奇数,0 关闭;
  • mask_threshold(默认 4,0–50):近白转白、近黑转黑的阈值,0 关闭;
  • learn_mask(默认 False):分配一部分模型容量去复现输入掩码,以更多 VRAM 换取对复杂掩码的复现能力。

注意 ModelBase.__init__ 中的前置校验与此呼应:penalized_mask_losslearn_mask 开启而 mask_typenone 时会直接报错退出。

9.6 Optimizer 节区

配置项 默认 范围/选项 fixed 说明要点
optimizer adam adabelief/adam/adamax/adamw/lion/nadam/rms-prop AdaBelief 建议把 Epsilon Exponent 调到约 -16;AdamW 默认应配 weight decay 0.004;Lion 的学习率通常取 AdamW 的 1/3–1/10,weight decay 取 3–10 倍
learning_rate 5e-5 1e-6–1e-4 过大导致崩溃、过小易陷局部极小
epsilon_exponent -7 -20–0 即 1e-7;NaN 无解时调大 epsilon 可提升稳定性;Lion 不使用
save_optimizer exit never/always/exit 保存优化器权重使模型文件约增大 3 倍;exit 仅在显式停止或达到目标迭代时保存,断电/OOM/NaN 等异常退出不会保存
gradient_clipping none autoclip/global_norm/norm/value/none 防 NaN 手段,代价是额外 VRAM;autoclip 按梯度分布动态调整
clipping_value 1.0 0–10 autoclip 下为百分位(1.0→第 10 百分位),其余模式下为裁剪阈值
autoclip_history 10000 0–100000(步长 1000) autoclip 分析的迭代窗口,0 表示全历史
weight_decay 0.0 0–1 0.0 表示不衰减
gradient_accumulation 1 1–100 >1 时每 N 次迭代用平均梯度更新一次,小批次降噪
ada_beta_1 / ada_beta_2 0.9 / 0.999 0–1 一/二阶矩指数衰减率,适用于 AdaBelief、Adam、Adamax、AdamW、Lion、nAdam
ada_amsgrad False AMSGrad 变体,适用于 AdaBelief、Adam、AdamW

9.7 Trainer 节区:数据加载与增强

trainer_config.py 提供 LoaderAugmentation 两个节区,通过 get_defaults() 自动并入 config.ini 的 trainer.loader / trainer.augmentation 节。增强节区开头有官方警告:默认值对 99.9% 的使用场景都够用,除非你非常清楚自己在做什么,否则不要动

Loader 节区:

配置项 默认 范围 说明
num_processes 4 0–32 磁盘数据加载/处理进程数,0 表示仅主进程
pre_fetch 2 1–10 每个 loader 预取并驻留 RAM 的批数,磁盘读取速度波动大时才需要调整

Augmentation 节区(评估/图像/颜色三类):

配置项 默认 范围 说明
preview_images 14 2–16 训练预览中每侧展示的样例脸数
mask_opacity 30 0–100 预览中掩码叠加不透明度,越低越透明
mask_color #ff0000 颜色选择器 预览掩码叠加的 RGB 十六进制色
zoom_amount 5 0–25 随机缩放百分比
rotation_range 10 0–25 随机旋转百分比
shift_range 5 0–25 水平/垂直随机平移百分比
flip_chance 50 0–75 随机水平翻转概率(--no-flip 时忽略)
color_lightness 30 0–75 随机明度变化百分比(--no-augment-color 时忽略)
color_ab 8 0–50 L*a*b* 空间 a/b 通道随机变化百分比
color_clahe_chance 50 0–75 做 CLAHE 的概率,fixed=False
color_clahe_max_size 4 1–8 CLAHE 网格尺寸上限(按训练图尺寸计算的乘数,0 到该值间随机取)

对应地,TrainConfig 中的 augment_color/flip 分别由 CLI 的 --no-augment-color/--no-flip 开关控制,与上表 "ignored if …" 的说明一致。

10. 一次训练的完整调用链

把以上模块串起来,faceswap.py train -A <face_a> -B <face_b> 的执行脉络(基于源码结构推断)是:

  1. faceswap.py_main()generate_configs() 生成配置文件,再由 TrainArgs 解析参数;
  2. 模型插件(继承 ModelBase)构造时加载 train_config,初始化 IO/State/Settings,state 文件驱动配置恢复与遗留迁移;
  3. build() 构建或加载模型;Inference 只在 convert 路径生效;
  4. TrainerBase 子类创建(单卡用 original、多卡用 distributed),register_loss()LossCollator 挂载为模型的 loss_func 子模块;
  5. 每个训练批:DataLoader 产出 inputs/targets/metatrain_batch() 前向(Keras 模型 training=True)→ 逐侧损失 → 汇总后 optimizer.backward + stepModelBase.add_history(loss) 累计损失历史 → State.increment_iterations()
  6. 每 N 次保存迭代 IO.save().keras(按 save_optimizer 策略注入 optimizer.pt)、State.save() 写 state、按平均损失下降决定 Backupsnapshot_interval > 0 时额外触发 IO.snapshot()

11. 测试与文档入口

12. 小结

Faceswap 的 train 包用一套非常克制的分层设计支撑全部模型插件:ModelBase 统一输入约定(A/B 双输入、input_shapecolor_order)与生命周期(build → compile → save/snapshot),Inference 负责训练-推理形态转换,IO 承担 .keras 存取、优化器状态注入、基于损失下降的备份与损坏自愈(Legacy/PatchKerasConfig 自动迁移),State 以 fixed/updatable 双轨语义保证「老模型按老配置恢复、可运行参数可热更新」,Settings 用两次构建 + config 重建实现 FP32/混合精度互切;trainer 侧则以 TrainerBase 抽象 + WrappedModel/DataParallel 实现单卡与多卡的同一损失聚合路径。理解了这条链路,无论是阅读某个具体模型插件(如 original、dfaker)、调试恢复训练失败,还是自定义新模型插件,都有了明确的落点。

登录后查看全文
热门项目推荐
相关项目推荐