首页
/ Faceswap lib.training 包源码解析:训练循环、损失加权、优化器封装与学习率调优的完整机制

Faceswap lib.training 包源码解析:训练循环、损失加权、优化器封装与学习率调优的完整机制

2026-09-06 16:03:49作者:幸俭卉

本文基于 Faceswap 仓库的 lib/training 包 API 文档(对应 docs/full/lib/training.rst)展开。该包是 Faceswap 模型训练的"执行引擎",封装了训练循环、损失函数编排、优化器与梯度处理、学习率预热/自动搜索、训练预览图生成与 TensorBoard 日志等全部支撑库。读完后,你将能理解一次 faceswap train 训练过程中各模块的调用关系、可配置参数的底层含义,以及如何从源码层面解释损失曲线、学习率扫描等训练行为。

1. 包总览:lib.training 在训练流程中的位置

lib/training/init.py 的模块文档字符串定义为:负责处理"对齐文件、检测到的脸、对齐后的脸及其相关对象",并在导入时决定预览后端:优先尝试 preview_tk.PreviewTk(Tkinter 桌面窗口),若导入失败则回退到 preview_cv.PreviewCV(OpenCV 窗口)。从源码结构看,该包由以下模块组成,与 API 文档中的 automodapi 列表一一对应:

模块 文件 职责
训练循环 lib/training/train.py 所有 Trainer 插件的基类,驱动逐批训练
损失计算 lib/training/loss.py 损失函数拼装、加权、掩码与眼部/嘴部加权
学习率搜索 lib/training/lr_finder.py 指数扫描学习率并绘制 loss 曲线
学习率预热 lib/training/lr_warmup.py Warmup 阶段线性提升学习率
优化器封装 lib/training/optimizer.py 优化器选择、梯度裁剪、混合精度、梯度累积
预览图生成 lib/training/preview.py 预测图/真值图对比拼版
预览窗口(CV) lib/training/preview_cv.py OpenCV 实时预览
预览窗口(Tk) lib/training/preview_tk.py Tkinter 预览与任务栏控件
TensorBoard lib/training/tensorboard.py 训练日志写入与读取
数据管线 lib/training/data/ 数据集、增强、拼批与加载器

2. Trainer:训练循环基类

Trainer 是"处理训练图像送入模型、生成 TensorBoard 日志、创建 sample/time-lapse 预览图"的核心,所有 Trainer 插件必须继承自它。构造参数(来自类 docstring,可直接对照源码):

  • plugin:负责处理每个 batch 的训练插件(具体网络结构由 plugins/train/trainer/ 下各插件实现);
  • preview:是否生成预览图;
  • warmup_steps:学习率预热步数,默认 0;
  • timelapse_folders:用于生成 timelapse 的输入文件夹,默认 None(不生成);
  • timelapse_output:timelapse 输出路径。

其内部方法体现了完整的一次训练步调用链:

  1. _get_train_loader() / _get_preview_loader() / _get_timelapse_loader() 分别构建训练、预览、timelapse 三个数据加载器(均基于 lib/training/data/loader.py);
  2. _handle_lr_finder():训练前可选地先跑一次学习率扫描(见第 4 节);
  3. train_one_batch():执行一次前向 + 损失计算,返回 list[BatchLoss]
  4. _get_predictions()_update_viewers():把 batch 预测转为预览图并推送到预览窗口;
  5. train_one_step():串起"取样 → 训练 → 更新预览 → 计时"的完整一步;
  6. save() / exit_early():保存模型与控制训练退出。

训练过程中捕获 torch.cuda.OutOfMemoryError 等异常(见文件头部导入),使显存不足时训练循环可以受控处理而非直接崩溃。

3. LossCollator 与 BatchLoss:损失函数的加权与掩码机制

LossCollator 是一个 nn.Module,负责"编译所选损失函数并在训练循环中计算数值"。其构造参数完整列出了 Faceswap 损失配置项的语义:

  • functions / weights:来自配置文件的损失函数名列表与对应权重列表,两者长度必须一致,否则抛出 ValueErrorloss.py#L145-L147);
  • color_order:模型训练使用的颜色通道顺序(bgr/rgb),部分损失函数(如 SSIM)对通道顺序敏感;
  • use_mask:是否启用 penalize mask loss,即仅在面部掩码区域内计算损失;
  • eye_multiplier / mouth_multiplier:对眼部/嘴部区域施加的额外权重倍数;
  • smallest_output:模型最小输出尺寸,某些损失函数初始化时需要;
  • mask_loss:启用 learn_mask 时使用的掩码损失函数,默认 None

3.1 空间损失与非空间损失的自动分类

_get_function_types()loss.py#L161-L188)的实现很有巧思:它会构造两个 (1, 3, size, size) 的随机张量,让每个损失函数跑一次前向,按输出维度自动分类——4 维 (N, C, H, W) 输出为"空间损失"(spatial),1 维 (N,) 标量输出为"非空间损失"(non-spatial),其他维度直接抛 RuntimeError。这决定了后续两条不同的计算路径:

  • 空间损失_get_spatial_lossloss.py#L190-L223):逐像素损失可被掩码直接乘上去——启用 use_maskloss *= mask_face;眼部/嘴部倍数大于 1 时执行 loss += loss * mask_eye * eye_multiplier,即在眼部区域把损失放大;
  • 非空间损失_get_non_spatial_lossloss.py#L266-L301):标量损失无法乘像素掩码,因此改为对输入先做掩码——_get_masked_inputs() 会分别构造"整脸掩码输入(权重 1.0)""眼部掩码输入(权重 eye_multiplier)""嘴部掩码输入(权重 mouth_multiplier)"三组,每组独立计算标量损失后加权求和。

3.2 BatchLoss 数据类

BatchLoss 是一个 dataclass,保存每个输出侧的 unweighted(未加权)与 weighted(加权后)逐函数损失字典,以及可选的 mask 掩码损失。total 属性按"对所有加权损失逐元素求和再取 mean"得到用于反向传播的标量;to_cpu() 会把全部张量 detach 并搬到 CPU 供打印和 TensorBoard 记录使用。这一设计使控制台 loss 输出与张量日志天然区分了"未加权原始值"和"实际参与梯度下降的加权值"。

4. 学习率调优:WarmupScheduler 与 LearningRateFinder

4.1 学习率预热(lr_warmup)

WarmupScheduler 继承 torch.optim.lr_scheduler.LRScheduler,核心逻辑在 get_lr() 中:当前步数未达到 steps 时,学习率按 base_lr * (last_epoch / steps) 线性从接近 0 爬到目标值;达到 steps 后直接返回 base_lrs 保持恒定。它还会在 10%、20%…100% 的"汇报点"(_reporting_points,构造函数中按 steps * i / 10 生成)打印进度日志,方便用户确认预热是否按预期完成。在 Optimizer.step() 中可以看到它的使用条件:_warmup 存在且 _session_steps < warmup.steps 时逐步 step,否则跳过。

4.2 学习率扫描(lr_finder)

LearningRateFinder 实现经典的学习率范围扫描(LR Range Test):

  • strength 档位:由枚举 LRStrength 定义——default=10aggressive=5extreme=2.5。该数值决定从"最陡下降点"往回退的倍率,数值越小选出的学习率越激进;
  • mode 参数set(扫描后直接设置学习率)、graph_and_set(画 loss 曲线图后设置)、graph_and_exit(画曲线图后退出,供人工判读);
  • 平滑与早停:每个 batch 结束后记录学习率与损失,用指数滑动平均(beta=0.98)做平滑;当平滑损失超过 stop_factor(=4) × 历史最优 时判定发散、提前退出,遇到 NaN 同样提前退出(见 _on_batch_end);
  • 调度机制:由 Optimizer.find_learning_rate() 驱动——先记录原始 lr 与优化器状态,用 ExponentialLR(gamma 由 (end_lr/start_lr)**(1/steps) 计算)在 steps 步内把学习率从 start_lr 指数增长到 end_lr,扫描结束后恢复优化器/scaler 状态并调用 set_lr(best_lr) 写入最优值。

训练插件侧通过 Trainer._handle_lr_finder() 决定训练开始前是否触发这一流程。

5. Optimizer 封装:优化器选择、参数分组与梯度处理

Optimizer 类"管理所选的 Torch 优化器",是训练配置项到 PyTorch API 的翻译层。

5.1 支持的优化器与超参映射

文件头部的 _OPTIMIZERS 映射表(optimizer.py#L30-L36)列出全部可选优化器:adabeliefadamadamaxadamwlionnadamrms-prop_get_optimizer_kwargs()optimizer.py#L181-L206)展示了用户配置项到 torch 参数的精确翻译规则:

配置项 torch 参数 适用优化器
weight_decay() weight_decay 全部(按参数分组应用,见下)
epsilon_exponent() eps = 10 ** epsilon_exponent 除 lion 外全部
ada_beta_1() / ada_beta_2() betas = (β1, β2) adabelief/adam/adamw/adamax/lion/nadam
ada_amsgrad() amsgrad adabelief/adam/adamw
learning_rate() lr 全部

5.2 权重衰减的参数分组

_get_parameter_groups()optimizer.py#L233-L271)配合 get_parameter_group_ids 把模型权重拆成两组:一维权重或名称以 bias 结尾的参数归入 no_decayweight_decay=0.0),其余归入 decay 组应用配置的衰减量。这是标准的"不对 bias/归一化参数做 L2 正则"做法,避免了归一化层参数被不当衰减。

5.3 梯度裁剪的四种方法

GradClip 支持 methodautoclipglobal_normnormvaluenone 时不实例化裁剪器):

  • autoclip:基于 AutoClipperlib/model/ 包),按历史梯度范数百分位自动决定裁剪阈值——value 表示百分位数(1.0 → 10 百分位,2.5 → 25 百分位),历史长度由 autoclip_history 控制(默认 10000 步);
  • global_norm:torch 原生 nn.utils.clip_grad_norm_,按全局范数裁剪;
  • norm:自定义 _clip_norm()optimizer.py#L82-L100),逐参数按各自的 L2 范数独立裁剪;
  • value:torch 原生 nn.utils.clip_grad_value_,按逐元素值裁剪。

5.4 梯度累积、混合精度与状态持久化

backward()optimizer.py#L357-L369)先把损失除以 accumulation_steps 再调用 backward(),实现"等效更大 batch"的梯度累积;启用 mixed_precision 时损失先经 GradScaler.scale() 缩放防止 fp16 下溢。step()optimizer.py#L371-L398)只在累积计数达到 accumulation_steps 时才真正更新:先 unscale_ 再裁剪,随后 optimizer.step()(混合精度下走 scaler.step() + scaler.update()),最后 zero_grad(set_to_none=True) 并归零计数。

state_dict() 序列化版本号为 1.0,包含优化器状态与 scaler 状态;_load_state() 在恢复训练时读取模型文件内保存的优化器状态,并支持从旧版 Keras 优化器状态迁移(版本 0.5 时走 _from_legacy(),校验参数个数、张量形状与参数组数量,不一致则重置优化器并给出警告)。set_lr() 会同步更新每个参数组的 lrinitial_lr,保证学习率扫描后旧调度器不会用错基准值。

6. 数据管线:data 子包四件套

API 文档的 "data package" 小节对应 lib/training/data/ 下的四个模块,构成"数据集 → 增强 → 拼批 → 加载器"的完整管线:

  • data_set.pylib/training/data/data_set.py):定义训练侧/预览侧数据集。其中 SideDataSetside(src/dst)、图像文件夹构造样本;_get_face() 负责按 sizecoverage 对齐人脸;Masks 相关方法(_get_landmarks_mask_get_face_mask)按 faceface_extendedeyemouth 类型生成掩码,供第 3 节的损失掩码机制消费;LabelledDataSet/UnlabelledDataSet 分别面向多身份训练与单身份场景(get_label() 提供身份标签);ConcatDataSet 可把多个数据集拼接并支持 shuffle。
  • augmentation.pylib/training/data/augmentation.py):训练期随机增强。ConstantsAugmentation.from_config()processing_size/batch_size 生成各增强操作的参数表;Augmentation 实现 random_flip()(随机水平翻转)、color_adjust()(LAB 空间随机色度抖动 + 随机 CLAHE 对比度增强)、warp()(基于 landmark 的仿射/非仿射形变,支持 to_landmarks 模式)。增强结果会同步变换 landmark,保证对齐一致。
  • collate.pylib/training/data/collate.py):把样本列表拼成训练 batch。LandmarkMatcher 通过 KD 树类逻辑(_get_closest_indices())为源/目标样本寻找 landmark 最接近的配对,支持 num_choices 控制候选数量;Collator.__call__() 执行 batch 内 resize、创建多尺度 targets,并产出 (inputs, targets, BatchMeta) 三元组——BatchMeta 携带每个输出尺寸对应的 face/eye/mouth 掩码,正是 LossCollator.forward()meta.mask_face[index] 等索引数据的来源。
  • loader.pylib/training/data/loader.py):TrainLoader 基于 torch.utils.data.DataLoader(支持 RandomSampler/DistributedSampler 注入,为分布式训练预留);PreviewLoader 则面向预览/timelapse,支持 num_samples 限制取样数量。两者都实现了迭代器协议(__iter__/__next__),get_loader() 返回实际 DataLoader 实例。

7. 预览与 TensorBoard:训练过程的可观测性

7.1 预览图拼版

lib/training/preview.py 中的 Samples 类按 coverage_ratiohas_maskmask_opacitymask_color 配置生成对比图:_get_background() 从 targets 抠出背景、_get_foreground() 把模型 predictions 叠加到原脸上,_apply_masks() 可选地渲染掩码区域;get_preview() 最终返回拼版后的 uint8 图像(_get_headers() 会为多路 swap 生成表头行)。

7.2 双后端预览窗口

  • preview_cv.pyPreviewCV 面向无桌面环境的 OpenCV 窗口,_should_shutdown() 检查窗口是否被用户关闭,可配置 TriggerType 按键触发行为;PreviewBuffer 作为线程间共享的图片缓冲(is_updated()/add_image()/get_images());
  • preview_tk.pyPreviewTk 提供 Tkinter 独立窗口或嵌入面板两种形态,任务栏(_Taskbar)含缩放 combo(_add_scale_combo)、插值方式 radio(_add_interpolator_radio)、保存按钮(_add_save_button);画布支持滚轮缩放(_on_bound_zoom)、拖拽平移(_on_mouse_drag)与键盘移动(_on_key_move),save_preview() 支持导出当前预览图。

Trainer._update_viewers() 将预览图回调注入这两类窗口,toggle_mask() 可动态切换掩码显示。

7.3 TensorBoard 日志

lib/training/tensorboard.py 包含两类实现:TensorBoard 是 Keras 回调风格(on_train_batch_endon_saveon_train_end 等钩子,update_freq 支持 batch/epoch/整数,write_graph 控制是否写入模型结构图);TorchTensorBoard 面向当前 torch 训练栈,配置项为 log_dir(默认 logs)、write_graph(默认 True)、update_freq(默认 epoch)。Trainer._set_tensorboard() 负责创建实例,_log_tensorboard() 把每步的 BatchLoss(各损失函数值、学习率等)写入日志,_clear_tensorboard() 在退出时清理。

8. 小结:一次训练步中的模块协作

把上述组件串起来,faceswap train 的一次训练步可以概括为:TrainLoaderdata_set 取样并施加 augmentationcollate 拼出带 BatchMeta 掩码的 batch;模型前向得到 predictions;LossCollator 按空间/非空间两条路径计算并加权损失(含眼部/嘴部倍数与面部掩码);Optimizer.backward()/step() 完成损失缩放、梯度累积、裁剪与参数更新,其间 WarmupScheduler 或 LRFinder 的 ExponentialLR 可能正在调整学习率;Trainer 把 predictions 交给 Samples 生成预览并推送到 Tk/OpenCV 窗口,同时把 loss 写入 TensorBoard;Trainer.save() 连同优化器 state_dict()(含 scaler 状态)一并持久化,保证断点续训时学习率、动量与梯度缩放状态都能精确恢复。这正是 docs/full/lib/training.rst 所索引的整套 lib.training 包的职责全貌。

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