Faceswap Train 插件包深度解析:Model 与 Trainer 的基类架构、状态管理与配置体系
本文以 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 包的两层结构:
- model 包:包含各模型插件可继承的辅助函数,覆盖六个基础模块——inference、io、model、settings、state、update,以及 original 这个带注释的模型插件范例。
- trainer 包:包含 Faceswap 的训练循环,即 base(抽象基类)、distributed(多卡训练)、original(标准训练器)与 trainer_config(数据加载与增强配置)。
从源码结构看,plugins/train/model/ 目录下还存在大量具体模型插件(dfaker.py、dlight.py、unbalanced.py、villain.py、phaze_a.py、realface.py、iae.py、lightweight.py、dfl_h128.py、dfl_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 命令行参数生成的参数命名空间;predict:True表示以推理(convert)方式加载模型,False表示以训练方式加载,默认False。
初始化时基类完成了若干关键动作(见 model.py#L45-L87):
- 设置
input_shape占位(模型插件必须在自己的__init__中覆写为 3 维 shape 元组;若 A/B 两侧输入尺寸不同,则应赋值为两个 shape 元组的列表); - 通过
cfg.load_config(config_file=...)加载 train_config.py 定义的全局训练配置; - 前置校验:若选择了
Penalized Mask Loss或Learn Mask却未选择任何 Mask(mask_type为none),直接抛出FaceswapError终止; - 实例化
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_sides:True表示 B→A 方向换脸,False表示 A→B 方向。
其内部流程(可对照 inference.py#L29-L205):
_get_input:断言模型输入数为 2,取side_idx(B→A 为 0,A→B 为 1)对应的输入张量;_get_valid_layer_inputs:对每个非InputLayer层,通过 Keras 内部_inbound_nodes找出该层在所选换脸方向下的有效输入层列表;_get_output_layer:从model.output的 Keras history 中找出输出层——支持两种结构:共享解码器(全部输出指向同一层)与拆分解码器(A/B 各自独立输出层,按side_idx取前半或后半);_backwards_recurse:从输出层反向递归收集属于该换脸路径的全部层(带seen集合防环);__call__:按拓扑依赖逐层调用已收集到的层实例,重建出一个名为{原模型名}_inference的单输入 Keras 推理模型返回。
这意味着 convert 阶段不需要加载整个双塔网络的全部权重路径,只裁剪出实际用到的半边,从而降低显存占用。
4. io 模块:模型存取、优化器注入与备份
io.py 由 IO、Weights、OptimizerMigrate 三个类与 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-weights:freeze()先统一解冻所有子模型(因为layer.trainable检查在冻结后仍可能误报True),再把插件freeze_layers(默认["encoder"])中实际存在的层设为trainable = False;找不到对应层名时输出警告;--load-weights:load()要求权重文件必须是有效.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_scale、dynamic_growth_steps、step_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:训练会话字典,每个会话记录timestamp、no_logs、batchsize、iterations、config(仅包含可更新项);lowest_avg_loss:历史最低保存间隔平均损失(用于 io 的备份判定);iterations:总迭代数;mixed_precision_layers:可混合精度切换的层名列表;lr_finder:学习率查找器发现的值(-1 表示无);config:完整训练配置快照。
5.1 fixed 与 updatable:配置项的双轨语义
恢复训练时,_update_config() 的核心逻辑是以 state 文件中的配置为权威:
fixed=True的配置项(如centering、coverage、optimizer):若用户当前配置与 state 不一致,用 state 中的值覆盖用户配置(opt.set(old_val)),即改这些项对已有模型无效;fixed=False的配置项(如learning_rate、loss_function、mixed_precision):允许运行时更新,新值写回 state,并记入_updatable_options;若变更项属于rebuild_tasks(目前是mixed_precision),置model_needs_rebuild = True触发模型重建;- 对 state 中缺失的新增配置项,优先采用
legacy_defaults(centering="legacy"、coverage=62.5、mask_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_type 为 facehull/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.py 的 Settings 在模型启动前完成后端设置。混合精度的启用条件是 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) 直接执行完整升级流程:
_get_model_config():绕过 Keras 3 的加载器,用h5py直接从.h5的 attrs 中读取keras_version与model_config,校验主版本必须为 2;_update_layers()递归处理每层配置:LeakyReLU的alpha参数改名为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 不接受嵌套输出);
_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 名/inboundkeras_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 是抽象基类:持有 model、batch_size、config、sampler(由子类 get_sampler() 提供,类型限定为 RandomSampler 或 DistributedSampler),register_loss() 会把配置的 LossCollator 以 add_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.py 的 Trainer 继承 original,核心改动:
WrappedModel:把双输入 Keras 模型包成单入口torch.nn.Module,其forward(inputs, targets, meta_dict)在每张卡上完成前向并计算各侧损失,返回损失字典列表(过滤 None 项);_validate_batch_size():批次小于 GPU 数时自动提升到 GPU 数并告警;批次不能被 GPU 数整除时提示更优值((batch // gpus) * gpus或其加一个 GPU 数)——因为DataParallel按卡均分批次;_set_distributed():用torch.nn.DataParallel(WrappedModel(...))包装模型(CUDA_VISIBLE_DEVICES已由-X命令行参数预先设置),并拦截 Torch 关于 GPU 显存不均衡的UserWarning、裁剪掉与 Faceswap 无关的后缀提示;_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.ini 中 global 与 trainer.* 节区)。文档顶部有一句总注记:除特别说明外,此处修改的取值仅对新建模型生效——这正是 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):候选为ffl、gmsd、l_inf_norm、laploss、logcosh、mae、mse、ms_ssim、pixel_gradient_diff、smooth_loss、ssim; - 第二/三/四损失(
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):可选none、components、extended(凸包外扩到额头),外加PluginLoader.get_available_extractors("mask")动态提供的 NN 掩码插件(文档说明中列出了bisenet-fp_face、bisenet-fp_head、vgg-clear、vgg-obstructed、unet-dfl、custom_face、custom_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_loss 或 learn_mask 开启而 mask_type 为 none 时会直接报错退出。
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 提供 Loader 与 Augmentation 两个节区,通过 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> 的执行脉络(基于源码结构推断)是:
- faceswap.py 的
_main()先generate_configs()生成配置文件,再由TrainArgs解析参数; - 模型插件(继承
ModelBase)构造时加载train_config,初始化IO/State/Settings,state 文件驱动配置恢复与遗留迁移; build()构建或加载模型;Inference只在 convert 路径生效;TrainerBase子类创建(单卡用 original、多卡用 distributed),register_loss()把LossCollator挂载为模型的loss_func子模块;- 每个训练批:DataLoader 产出
inputs/targets/meta→train_batch()前向(Keras 模型training=True)→ 逐侧损失 → 汇总后optimizer.backward + step→ModelBase.add_history(loss)累计损失历史 →State.increment_iterations(); - 每 N 次保存迭代
IO.save()写.keras(按save_optimizer策略注入optimizer.pt)、State.save()写 state、按平均损失下降决定Backup;snapshot_interval > 0时额外触发IO.snapshot()。
11. 测试与文档入口
- 训练器行为有专门测试覆盖:tests/plugins/train/trainer/test_original.py 与 test_distributed.py,验证前向/反向与 DataParallel 包装路径;
- 插件加载机制的文档见 plugin_loader.rst,train 插件包文档的完整索引见 docs/full/plugins/plugins.rst;
- 本文引用的核心源码路径汇总:plugins/train/model/_base/model.py、plugins/train/model/_base/inference.py、plugins/train/model/_base/io.py、plugins/train/model/_base/settings.py、plugins/train/model/_base/state.py、plugins/train/model/_base/update.py、plugins/train/trainer/base.py、plugins/train/trainer/original.py、plugins/train/trainer/distributed.py、plugins/train/train_config.py、plugins/train/trainer/trainer_config.py。
12. 小结
Faceswap 的 train 包用一套非常克制的分层设计支撑全部模型插件:ModelBase 统一输入约定(A/B 双输入、input_shape、color_order)与生命周期(build → compile → save/snapshot),Inference 负责训练-推理形态转换,IO 承担 .keras 存取、优化器状态注入、基于损失下降的备份与损坏自愈(Legacy/PatchKerasConfig 自动迁移),State 以 fixed/updatable 双轨语义保证「老模型按老配置恢复、可运行参数可热更新」,Settings 用两次构建 + config 重建实现 FP32/混合精度互切;trainer 侧则以 TrainerBase 抽象 + WrappedModel/DataParallel 实现单卡与多卡的同一损失聚合路径。理解了这条链路,无论是阅读某个具体模型插件(如 original、dfaker)、调试恢复训练失败,还是自定义新模型插件,都有了明确的落点。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0624
Hy4-previewHy4 preview 是由腾讯混元团队研发的新一代混合专家(MoE)旗舰模型。模型总参数量 770B,每个 token 激活 49B,主干共包含78层,第一层采用标准 FFN,其余 77 层均为 MoE 结构,每层包含 256 个路由专家与 1 个共享专家,每个 token 激活 top-8 路由专家及共享专家。主干之外原生内置 1 层 MTP(总参数量 10B,激活 0.7B)以支持投机解码。Python00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
GLM-5.3-FlashGLM-5.3-Flash (320B-A18B),是GLM-5系列的首个原生多模态模型。320B总参数,能力超过GLM-5.2Jinja00
Spark-X2.5-4BSpark-X2.5-4B 旨在让强大的 AI 更实用、更高效、更易获得。在广泛日常任务中表现强劲,涵盖对话、写作、翻译、推理、编码、工具调用以及智能体工作流,并在同等规模的开源模型中取得领先成绩。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00